From 0f92cc9a6a4d856e6b6ce1895632da55034859aa Mon Sep 17 00:00:00 2001 From: cppla Date: Sat, 22 Aug 2026 22:06:43 +0800 Subject: [PATCH 1/2] feat: harden dual-ended acceleration stack --- .github/dependabot.yml | 29 + .github/workflows/ci.yml | 78 +- .github/workflows/codeql.yml | 8 +- .github/workflows/netem.yml | 6 +- Dockerfile | 9 +- Makefile | 46 +- README.md | 190 +- SECURITY.md | 86 +- THIRD_PARTY_NOTICES.md | 1019 +++++ cmd/autocar/bench.go | 27 +- cmd/autocar/common.go | 153 +- cmd/autocar/main.go | 2 +- cmd/autocar/main_test.go | 46 + cmd/autocar/server.go | 136 +- docs/ACCELERATION.md | 195 + docs/ARCHITECTURE.md | 211 +- docs/BENCHMARK.md | 223 +- docs/DEPLOYMENT.md | 530 ++- docs/PROTOCOL.md | 249 +- docs/assets/autocar-logo.svg | 99 + go.mod | 28 +- go.sum | 31 +- internal/hy2/auto.go | 281 ++ internal/hy2/client.go | 543 +++ internal/hy2/hy2_test.go | 1614 ++++++++ internal/hy2/server.go | 622 +++ internal/hy2/tls_policy.go | 84 + internal/proxy/config.go | 3 + internal/proxy/socks5.go | 471 ++- internal/proxy/socks5_udp_test.go | 642 +++ internal/security/dialer.go | 62 + internal/security/dialer_test.go | 100 + internal/transport/transport.go | 24 + internal/tunnel/common.go | 9 +- internal/tunnel/tls.go | 93 +- internal/tunnel/tls_limits_test.go | 233 ++ scripts/check-fork-provenance.sh | 47 + scripts/govulncheck.sh | 45 + scripts/netem-integration.sh | 265 +- third_party/hysteria-core/AUTOCAR_PATCHES.md | 51 + third_party/hysteria-core/LICENSE.md | 7 + .../hysteria-core/client/.mockery.yaml | 9 + third_party/hysteria-core/client/client.go | 383 ++ third_party/hysteria-core/client/config.go | 136 + .../hysteria-core/client/fast_open_test.go | 87 + .../hysteria-core/client/mock_udpIO.go | 139 + third_party/hysteria-core/client/reconnect.go | 120 + third_party/hysteria-core/client/udp.go | 188 + third_party/hysteria-core/client/udp_test.go | 146 + third_party/hysteria-core/errors/errors.go | 75 + third_party/hysteria-core/go.mod | 32 + third_party/hysteria-core/go.sum | 46 + .../internal/congestion/bbr/bandwidth.go | 27 + .../congestion/bbr/bandwidth_sampler.go | 877 +++++ .../internal/congestion/bbr/bbr_sender.go | 1088 +++++ .../congestion/bbr/bbr_sender_test.go | 226 ++ .../internal/congestion/bbr/clock.go | 18 + .../bbr/packet_number_indexed_queue.go | 199 + .../internal/congestion/bbr/ringbuffer.go | 118 + .../congestion/bbr/windowed_filter.go | 162 + .../internal/congestion/brutal/brutal.go | 193 + .../internal/congestion/brutal/brutal_test.go | 45 + .../internal/congestion/common/pacer.go | 80 + .../internal/congestion/utils.go | 72 + .../hysteria-core/internal/frag/frag.go | 110 + .../hysteria-core/internal/frag/frag_test.go | 385 ++ .../internal/integration_tests/.mockery.yaml | 29 + .../integration_tests/chrome_parrot_test.go | 70 + .../internal/integration_tests/close_test.go | 252 ++ .../internal/integration_tests/hook_test.go | 146 + .../internal/integration_tests/masq_test.go | 93 + .../mocks/mock_Authenticator.go | 94 + .../integration_tests/mocks/mock_Conn.go | 426 ++ .../mocks/mock_EventLogger.go | 249 ++ .../integration_tests/mocks/mock_Outbound.go | 199 + .../mocks/mock_RequestHook.go | 188 + .../mocks/mock_TrafficLogger.go | 184 + .../integration_tests/mocks/mock_UDPConn.go | 197 + .../internal/integration_tests/smoke_test.go | 283 ++ .../internal/integration_tests/stress_test.go | 263 ++ .../integration_tests/trafficlogger_test.go | 180 + .../integration_tests/udp_acl_test.go | 177 + .../internal/integration_tests/utils_test.go | 104 + .../hysteria-core/internal/pmtud/avail.go | 7 + .../hysteria-core/internal/pmtud/unavail.go | 13 + .../hysteria-core/internal/protocol/http.go | 68 + .../internal/protocol/padding.go | 31 + .../hysteria-core/internal/protocol/proxy.go | 261 ++ .../internal/protocol/proxy_test.go | 330 ++ .../hysteria-core/internal/utils/atomic.go | 54 + .../hysteria-core/internal/utils/qstream.go | 62 + .../hysteria-core/server/.mockery.yaml | 15 + third_party/hysteria-core/server/config.go | 404 ++ third_party/hysteria-core/server/copy.go | 80 + .../server/copy_benchmark_test.go | 23 + .../hysteria-core/server/mock_UDPConn.go | 197 + .../server/mock_udpEventLogger.go | 100 + .../hysteria-core/server/mock_udpIO.go | 290 ++ .../server/resource_limits_test.go | 409 ++ third_party/hysteria-core/server/server.go | 601 +++ third_party/hysteria-core/server/udp.go | 470 +++ third_party/hysteria-core/server/udp_test.go | 191 + .../quic-go/.clusterfuzzlite/Dockerfile | 5 + third_party/quic-go/.clusterfuzzlite/build.sh | 38 + .../quic-go/.clusterfuzzlite/project.yaml | 1 + third_party/quic-go/.githooks/README.md | 8 + third_party/quic-go/.githooks/pre-commit | 34 + third_party/quic-go/.github/FUNDING.yml | 13 + third_party/quic-go/.github/dependabot.yml | 6 + .../workflows/build-interop-docker.yml | 51 + .../workflows/clusterfuzz-coverage.yml | 106 + .../workflows/clusterfuzz-lite-batch.yml | 87 + .../.github/workflows/clusterfuzz-lite-pr.yml | 51 + .../workflows/clusterfuzz-lite-prune.yml | 29 + .../quic-go/.github/workflows/codspeed.yml | 30 + .../.github/workflows/cross-compile.sh | 33 + .../.github/workflows/cross-compile.yml | 50 + .../quic-go/.github/workflows/go-generate.sh | 18 + .../quic-go/.github/workflows/govulncheck.yml | 19 + .../quic-go/.github/workflows/integration.yml | 89 + .../quic-go/.github/workflows/lint.yml | 101 + .../quic-go/.github/workflows/unit.yml | 75 + third_party/quic-go/.gitignore | 20 + third_party/quic-go/.golangci.yml | 99 + third_party/quic-go/AUTOCAR_PATCHES.md | 28 + third_party/quic-go/FIPS140.md | 37 + third_party/quic-go/FUZZING.md | 59 + third_party/quic-go/LICENSE | 21 + third_party/quic-go/README.md | 65 + third_party/quic-go/SECURITY.md | 14 + third_party/quic-go/assets/LICENSE.md | 27 + third_party/quic-go/assets/logo.svg | 330 ++ third_party/quic-go/assets/quic-go-logo.png | Bin 0 -> 105123 bytes third_party/quic-go/buffer_pool.go | 92 + third_party/quic-go/buffer_pool_test.go | 44 + third_party/quic-go/chrome_parrot.go | 59 + third_party/quic-go/client.go | 109 + third_party/quic-go/client_test.go | 105 + third_party/quic-go/closed_conn.go | 58 + third_party/quic-go/closed_conn_test.go | 34 + third_party/quic-go/codecov.yml | 23 + third_party/quic-go/config.go | 153 + third_party/quic-go/config_test.go | 250 ++ third_party/quic-go/congestion/interface.go | 67 + third_party/quic-go/conn_id_generator.go | 212 + third_party/quic-go/conn_id_generator_test.go | 350 ++ third_party/quic-go/conn_id_manager.go | 321 ++ third_party/quic-go/conn_id_manager_test.go | 432 ++ third_party/quic-go/conn_wrapped_test.go | 73 + third_party/quic-go/connection.go | 3214 +++++++++++++++ third_party/quic-go/connection_logging.go | 315 ++ .../quic-go/connection_logging_test.go | 143 + third_party/quic-go/connection_test.go | 3485 +++++++++++++++++ third_party/quic-go/crypto_stream.go | 283 ++ third_party/quic-go/crypto_stream_manager.go | 73 + .../quic-go/crypto_stream_manager_test.go | 87 + third_party/quic-go/crypto_stream_test.go | 288 ++ third_party/quic-go/datagram_queue.go | 137 + third_party/quic-go/datagram_queue_test.go | 180 + third_party/quic-go/errors.go | 107 + third_party/quic-go/errors_test.go | 37 + third_party/quic-go/example/client/main.go | 81 + third_party/quic-go/example/echo/echo.go | 113 + third_party/quic-go/example/main.go | 182 + third_party/quic-go/flow_controller_base.go | 84 + .../quic-go/flow_controller_connection.go | 186 + .../flow_controller_connection_test.go | 77 + third_party/quic-go/flow_controller_stream.go | 196 + .../quic-go/flow_controller_stream_test.go | 319 ++ .../flow_controller_test_helpers_test.go | 39 + third_party/quic-go/frame_sorter.go | 256 ++ third_party/quic-go/frame_sorter_test.go | 1653 ++++++++ third_party/quic-go/framer.go | 295 ++ third_party/quic-go/framer_test.go | 476 +++ third_party/quic-go/go.mod | 35 + third_party/quic-go/go.sum | 107 + third_party/quic-go/http3/README.md | 9 + third_party/quic-go/http3/body.go | 137 + third_party/quic-go/http3/body_test.go | 140 + third_party/quic-go/http3/capsule.go | 148 + third_party/quic-go/http3/capsule_test.go | 163 + third_party/quic-go/http3/client.go | 496 +++ third_party/quic-go/http3/client_test.go | 887 +++++ third_party/quic-go/http3/conn.go | 319 ++ third_party/quic-go/http3/conn_test.go | 502 +++ third_party/quic-go/http3/error.go | 63 + third_party/quic-go/http3/error_codes.go | 84 + third_party/quic-go/http3/error_codes_test.go | 37 + third_party/quic-go/http3/error_test.go | 87 + third_party/quic-go/http3/frames.go | 327 ++ third_party/quic-go/http3/frames_test.go | 487 +++ third_party/quic-go/http3/gzip_reader.go | 39 + third_party/quic-go/http3/headers.go | 429 ++ third_party/quic-go/http3/headers_test.go | 776 ++++ .../quic-go/http3/http3_helper_test.go | 353 ++ .../quic-go/http3/internal/testdata/cert.go | 28 + .../http3/internal/testdata/cert_test.go | 31 + third_party/quic-go/http3/ip_addr.go | 48 + .../quic-go/http3/mock_clientconn_test.go | 157 + .../http3/mock_datagram_stream_test.go | 536 +++ .../quic-go/http3/mock_quic_listener_test.go | 158 + third_party/quic-go/http3/mockgen.go | 11 + third_party/quic-go/http3/qlog.go | 56 + third_party/quic-go/http3/qlog/event.go | 138 + third_party/quic-go/http3/qlog/event_test.go | 98 + third_party/quic-go/http3/qlog/frame.go | 220 ++ third_party/quic-go/http3/qlog/frame_test.go | 204 + third_party/quic-go/http3/qlog/qlog_dir.go | 15 + .../quic-go/http3/qlog/qlog_dir_test.go | 41 + third_party/quic-go/http3/request_writer.go | 324 ++ .../quic-go/http3/request_writer_test.go | 175 + third_party/quic-go/http3/response_writer.go | 372 ++ .../quic-go/http3/response_writer_test.go | 256 ++ third_party/quic-go/http3/server.go | 782 ++++ third_party/quic-go/http3/server_conn.go | 259 ++ third_party/quic-go/http3/server_test.go | 885 +++++ .../quic-go/http3/state_tracking_stream.go | 173 + .../http3/state_tracking_stream_test.go | 319 ++ third_party/quic-go/http3/stream.go | 406 ++ third_party/quic-go/http3/stream_test.go | 235 ++ third_party/quic-go/http3/trace.go | 105 + third_party/quic-go/http3/transport.go | 538 +++ third_party/quic-go/http3/transport_test.go | 572 +++ .../integrationtests/self/benchmark_test.go | 151 + .../integrationtests/self/cancelation_test.go | 598 +++ .../self/chrome_parrot_test.go | 195 + .../integrationtests/self/close_test.go | 230 ++ .../integrationtests/self/conn_id_test.go | 151 + .../self/connection_migration_test.go | 144 + .../integrationtests/self/datagram_test.go | 432 ++ .../integrationtests/self/deadline_test.go | 235 ++ .../integrationtests/self/drop_test.go | 117 + .../integrationtests/self/early_data_test.go | 72 + .../self/handshake_context_test.go | 289 ++ .../self/handshake_drop_test.go | 411 ++ .../self/handshake_rtt_test.go | 201 + .../integrationtests/self/handshake_test.go | 820 ++++ .../self/http_datagram_test.go | 323 ++ .../self/http_hotswap_test.go | 111 + .../integrationtests/self/http_qlog_test.go | 83 + .../self/http_raw_conn_test.go | 171 + .../self/http_shutdown_test.go | 520 +++ .../integrationtests/self/http_test.go | 1426 +++++++ .../integrationtests/self/http_trace_test.go | 137 + .../integrationtests/self/key_update_test.go | 93 + .../integrationtests/self/mitm_test.go | 428 ++ .../quic-go/integrationtests/self/mtu_test.go | 197 + .../integrationtests/self/multiplex_test.go | 346 ++ .../self/nat_rebinding_test.go | 126 + .../self/packetization_test.go | 273 ++ .../integrationtests/self/qlog_dir_test.go | 75 + .../integrationtests/self/qlog_test.go | 138 + .../integrationtests/self/resumption_test.go | 149 + .../quic-go/integrationtests/self/rtt_test.go | 183 + .../integrationtests/self/self_go124_test.go | 9 + .../integrationtests/self/self_go125_test.go | 9 + .../self/self_suite_linux_test.go | 21 + .../self/self_suite_others_test.go | 7 + .../integrationtests/self/self_test.go | 352 ++ .../self/simnet_helper_test.go | 105 + .../self/stateless_reset_test.go | 127 + .../integrationtests/self/stream_test.go | 375 ++ .../integrationtests/self/timeout_test.go | 504 +++ .../integrationtests/self/zero_rtt_test.go | 1154 ++++++ .../quic-go/integrationtests/tools/crypto.go | 127 + .../integrationtests/tools/crypto_test.go | 99 + .../integrationtests/tools/israce/norace.go | 6 + .../integrationtests/tools/israce/race.go | 6 + .../integrationtests/tools/proxy/proxy.go | 372 ++ .../tools/proxy/proxy_test.go | 503 +++ .../quic-go/integrationtests/tools/qlog.go | 55 + .../versionnegotiation/handshake_test.go | 192 + .../versionnegotiation/rtt_test.go | 58 + .../versionnegotiation/test_helper_test.go | 111 + third_party/quic-go/interface.go | 243 ++ .../internal/ackhandler/ack_eliciting.go | 33 + .../internal/ackhandler/ack_eliciting_test.go | 74 + .../quic-go/internal/ackhandler/cc_adapter.go | 62 + .../internal/ackhandler/cc_adapter_ex.go | 69 + .../quic-go/internal/ackhandler/ecn.go | 340 ++ .../quic-go/internal/ackhandler/ecn_test.go | 353 ++ .../quic-go/internal/ackhandler/frame.go | 21 + .../quic-go/internal/ackhandler/interfaces.go | 45 + .../ackhandler/lost_packet_tracker.go | 73 + .../ackhandler/lost_packet_tracker_test.go | 75 + .../ackhandler/mock_ecn_handler_test.go | 189 + .../quic-go/internal/ackhandler/mockgen.go | 6 + .../quic-go/internal/ackhandler/packet.go | 60 + .../ackhandler/packet_number_generator.go | 84 + .../packet_number_generator_test.go | 92 + .../ackhandler/received_packet_handler.go | 119 + .../received_packet_handler_test.go | 144 + .../ackhandler/received_packet_history.go | 159 + .../received_packet_history_test.go | 304 ++ .../ackhandler/received_packet_tracker.go | 228 ++ .../received_packet_tracker_test.go | 188 + .../quic-go/internal/ackhandler/send_mode.go | 46 + .../internal/ackhandler/send_mode_test.go | 18 + .../ackhandler/sent_packet_handler.go | 1240 ++++++ .../ackhandler/sent_packet_handler_test.go | 1793 +++++++++ .../ackhandler/sent_packet_history.go | 274 ++ .../ackhandler/sent_packet_history_test.go | 343 ++ .../quic-go/internal/congestion/bandwidth.go | 22 + .../internal/congestion/bandwidth_test.go | 12 + .../quic-go/internal/congestion/clock.go | 20 + .../quic-go/internal/congestion/cubic.go | 214 + .../internal/congestion/cubic_sender.go | 330 ++ .../internal/congestion/cubic_sender_test.go | 594 +++ .../quic-go/internal/congestion/cubic_test.go | 205 + .../internal/congestion/hybrid_slow_start.go | 112 + .../congestion/hybrid_slow_start_test.go | 68 + .../quic-go/internal/congestion/interface.go | 33 + .../quic-go/internal/congestion/pacer.go | 110 + .../quic-go/internal/congestion/pacer_test.go | 154 + .../quic-go/internal/handshake/aead.go | 91 + .../quic-go/internal/handshake/aead_test.go | 108 + .../internal/handshake/chrome_client_hello.go | 76 + .../internal/handshake/cipher_suite.go | 114 + .../handshake/cipher_suite_fips140.go | 48 + .../internal/handshake/crypto_setup.go | 730 ++++ .../internal/handshake/crypto_setup_test.go | 574 +++ .../quic-go/internal/handshake/fake_conn.go | 21 + .../internal/handshake/fips140_go126.go | 9 + .../internal/handshake/fips140_legacy.go | 7 + .../internal/handshake/handshake_fuzz_test.go | 399 ++ .../handshake/handshake_helpers_test.go | 41 + .../internal/handshake/header_protector.go | 134 + .../quic-go/internal/handshake/hkdf.go | 26 + .../quic-go/internal/handshake/hkdf_test.go | 76 + .../internal/handshake/initial_aead.go | 80 + .../internal/handshake/initial_aead_test.go | 317 ++ .../quic-go/internal/handshake/interface.go | 140 + .../internal/handshake/quic_event_go125.go | 11 + .../internal/handshake/quic_event_go126.go | 11 + .../quic-go/internal/handshake/retry_go125.go | 68 + .../quic-go/internal/handshake/retry_go126.go | 70 + .../quic-go/internal/handshake/retry_test.go | 56 + .../internal/handshake/session_ticket.go | 55 + .../internal/handshake/session_ticket_test.go | 46 + .../internal/handshake/tls_config_go126.go | 54 + .../handshake/tls_config_go126_test.go | 99 + .../internal/handshake/tls_config_go127.go | 24 + .../quic-go/internal/handshake/tls_conn.go | 27 + .../internal/handshake/tls_conn_utls.go | 266 ++ .../internal/handshake/token_generator.go | 126 + .../handshake/token_generator_test.go | 140 + .../internal/handshake/token_protector.go | 78 + .../handshake/token_protector_test.go | 70 + .../internal/handshake/updatable_aead.go | 372 ++ .../internal/handshake/updatable_aead_test.go | 736 ++++ .../mocks/ackhandler/sent_packet_handler.go | 713 ++++ .../quic-go/internal/mocks/congestion.go | 486 +++ .../quic-go/internal/mocks/crypto_setup.go | 730 ++++ .../internal/mocks/long_header_opener.go | 154 + third_party/quic-go/internal/mocks/mockgen.go | 10 + .../internal/mocks/short_header_opener.go | 155 + .../internal/mocks/short_header_sealer.go | 191 + third_party/quic-go/internal/monotime/time.go | 90 + .../quic-go/internal/monotime/time_test.go | 78 + .../internal/protocol/connection_id.go | 127 + .../internal/protocol/connection_id_test.go | 92 + .../internal/protocol/encryption_level.go | 65 + .../protocol/encryption_level_test.go | 40 + .../quic-go/internal/protocol/key_phase.go | 36 + .../internal/protocol/key_phase_test.go | 26 + .../internal/protocol/packet_number.go | 84 + .../internal/protocol/packet_number_test.go | 81 + .../quic-go/internal/protocol/params.go | 169 + .../quic-go/internal/protocol/params_test.go | 13 + .../quic-go/internal/protocol/perspective.go | 26 + .../internal/protocol/perspective_test.go | 18 + .../quic-go/internal/protocol/protocol.go | 156 + .../internal/protocol/protocol_test.go | 39 + .../quic-go/internal/protocol/stream.go | 102 + .../quic-go/internal/protocol/stream_test.go | 66 + .../quic-go/internal/protocol/version.go | 115 + .../quic-go/internal/protocol/version_test.go | 151 + .../quic-go/internal/qerr/error_codes.go | 87 + .../quic-go/internal/qerr/errorcodes_test.go | 47 + third_party/quic-go/internal/qerr/errors.go | 134 + .../quic-go/internal/qerr/errors_test.go | 176 + .../quic-go/internal/qtls/cipher_suite.go | 52 + .../internal/qtls/cipher_suite_test.go | 54 + third_party/quic-go/internal/testdata/cert.go | 143 + .../quic-go/internal/testdata/cert_test.go | 31 + .../internal/utils/buffered_write_closer.go | 26 + .../utils/buffered_write_closer_test.go | 25 + .../quic-go/internal/utils/connstats.go | 14 + .../internal/utils/linkedlist/README.md | 6 + .../internal/utils/linkedlist/linkedlist.go | 264 ++ third_party/quic-go/internal/utils/log.go | 131 + .../quic-go/internal/utils/log_test.go | 146 + third_party/quic-go/internal/utils/rand.go | 29 + .../quic-go/internal/utils/rand_test.go | 32 + .../internal/utils/ringbuffer/ringbuffer.go | 96 + .../utils/ringbuffer/ringbuffer_bench_test.go | 14 + .../utils/ringbuffer/ringbuffer_test.go | 49 + .../quic-go/internal/utils/rtt_stats.go | 159 + .../quic-go/internal/utils/rtt_stats_test.go | 146 + .../internal/utils/streamframe_interval.go | 45 + .../quic-go/internal/utils/tree/tree.go | 503 +++ .../internal/utils/tree/tree_match_test.go | 95 + .../quic-go/internal/utils/tree/tree_test.go | 254 ++ .../quic-go/internal/wire/ack_frame.go | 298 ++ .../quic-go/internal/wire/ack_frame_test.go | 612 +++ .../internal/wire/ack_frequency_frame.go | 65 + .../internal/wire/ack_frequency_frame_test.go | 71 + .../quic-go/internal/wire/ack_range.go | 14 + .../quic-go/internal/wire/ack_range_test.go | 12 + .../internal/wire/connection_close_frame.go | 75 + .../wire/connection_close_frame_test.go | 137 + .../quic-go/internal/wire/crypto_frame.go | 97 + .../internal/wire/crypto_frame_test.go | 123 + .../internal/wire/data_blocked_frame.go | 29 + .../internal/wire/data_blocked_frame_test.go | 40 + .../quic-go/internal/wire/datagram_frame.go | 85 + .../internal/wire/datagram_frame_test.go | 126 + .../quic-go/internal/wire/extended_header.go | 164 + .../internal/wire/extended_header_test.go | 269 ++ third_party/quic-go/internal/wire/frame.go | 33 + .../quic-go/internal/wire/frame_parser.go | 192 + .../internal/wire/frame_parser_test.go | 1083 +++++ .../quic-go/internal/wire/frame_test.go | 42 + .../quic-go/internal/wire/frame_type.go | 81 + .../quic-go/internal/wire/frame_type_test.go | 29 + .../internal/wire/handshake_done_frame.go | 17 + .../wire/handshake_done_frame_test.go | 16 + third_party/quic-go/internal/wire/header.go | 302 ++ .../quic-go/internal/wire/header_test.go | 793 ++++ .../internal/wire/immediate_ack_frame.go | 18 + .../internal/wire/immediate_ack_frame_test.go | 22 + third_party/quic-go/internal/wire/log.go | 74 + third_party/quic-go/internal/wire/log_test.go | 188 + .../quic-go/internal/wire/max_data_frame.go | 33 + .../internal/wire/max_data_frame_test.go | 39 + .../internal/wire/max_stream_data_frame.go | 43 + .../wire/max_stream_data_frame_test.go | 46 + .../internal/wire/max_streams_frame.go | 50 + .../internal/wire/max_streams_frame_test.go | 117 + .../internal/wire/new_connection_id_frame.go | 80 + .../wire/new_connection_id_frame_test.go | 88 + .../quic-go/internal/wire/new_token_frame.go | 43 + .../internal/wire/new_token_frame_test.go | 51 + .../internal/wire/path_challenge_frame.go | 32 + .../wire/path_challenge_frame_test.go | 37 + .../internal/wire/path_response_frame.go | 32 + .../internal/wire/path_response_frame_test.go | 37 + .../quic-go/internal/wire/ping_frame.go | 17 + .../quic-go/internal/wire/ping_frame_test.go | 17 + third_party/quic-go/internal/wire/pool.go | 33 + .../quic-go/internal/wire/pool_test.go | 24 + .../internal/wire/reset_stream_frame.go | 79 + .../internal/wire/reset_stream_frame_test.go | 104 + .../wire/retire_connection_id_frame.go | 30 + .../wire/retire_connection_id_frame_test.go | 39 + .../quic-go/internal/wire/short_header.go | 62 + .../internal/wire/short_header_test.go | 106 + .../internal/wire/stop_sending_frame.go | 45 + .../internal/wire/stop_sending_frame_test.go | 47 + .../wire/stream_data_blocked_frame.go | 42 + .../wire/stream_data_blocked_frame_test.go | 46 + .../quic-go/internal/wire/stream_frame.go | 191 + .../internal/wire/stream_frame_test.go | 379 ++ .../internal/wire/streams_blocked_frame.go | 50 + .../wire/streams_blocked_frame_test.go | 117 + .../internal/wire/test_helpers_test.go | 32 + .../internal/wire/transport_parameter_test.go | 1172 ++++++ .../internal/wire/transport_parameters.go | 601 +++ .../wire/transport_parameters_chrome.go | 156 + .../wire/transport_parameters_chrome_test.go | 247 ++ .../internal/wire/version_negotiation.go | 53 + .../internal/wire/version_negotiation_test.go | 103 + third_party/quic-go/interop/Dockerfile | 38 + third_party/quic-go/interop/client/main.go | 210 + third_party/quic-go/interop/http09/client.go | 159 + .../quic-go/interop/http09/http_test.go | 77 + third_party/quic-go/interop/http09/server.go | 121 + third_party/quic-go/interop/run_endpoint.sh | 19 + third_party/quic-go/interop/server/main.go | 105 + third_party/quic-go/interop/utils/logging.go | 58 + .../quic-go/metrics/dashboards/README.md | 24 + .../metrics/dashboards/datasources.yml | 13 + .../metrics/dashboards/docker-compose.yml | 25 + .../quic-go/metrics/dashboards/prometheus.yml | 9 + .../quic-go/metrics/dashboards/quic-go.json | 926 +++++ .../quic-go/mock_ack_frame_source_test.go | 81 + third_party/quic-go/mock_conn_runner_test.go | 224 ++ third_party/quic-go/mock_frame_source_test.go | 121 + .../quic-go/mock_mtu_discoverer_test.go | 230 ++ third_party/quic-go/mock_packer_test.go | 395 ++ .../quic-go/mock_packet_handler_test.go | 149 + third_party/quic-go/mock_packetconn_test.go | 311 ++ third_party/quic-go/mock_raw_conn_test.go | 273 ++ .../quic-go/mock_sealing_manager_test.go | 197 + third_party/quic-go/mock_send_conn_test.go | 342 ++ third_party/quic-go/mock_sender_test.go | 264 ++ .../mock_stream_control_frame_getter_test.go | 82 + .../quic-go/mock_stream_frame_getter_test.go | 83 + .../quic-go/mock_stream_sender_test.go | 185 + third_party/quic-go/mock_unpacker_test.go | 124 + third_party/quic-go/mockgen.go | 47 + third_party/quic-go/module_rename.sh | 35 + third_party/quic-go/monotime/time.go | 37 + third_party/quic-go/mtu_discoverer.go | 253 ++ third_party/quic-go/mtu_discoverer_test.go | 242 ++ third_party/quic-go/oss-fuzz.sh | 52 + third_party/quic-go/packet_packer.go | 1111 ++++++ third_party/quic-go/packet_packer_chaos.go | 274 ++ .../quic-go/packet_packer_chaos_test.go | 292 ++ third_party/quic-go/packet_packer_test.go | 1086 +++++ third_party/quic-go/packet_unpacker.go | 222 ++ third_party/quic-go/packet_unpacker_test.go | 373 ++ third_party/quic-go/path_manager.go | 206 + third_party/quic-go/path_manager_outgoing.go | 314 ++ .../quic-go/path_manager_outgoing_test.go | 287 ++ third_party/quic-go/path_manager_test.go | 356 ++ third_party/quic-go/qlog/benchmark_test.go | 86 + third_party/quic-go/qlog/event.go | 849 ++++ third_party/quic-go/qlog/event_test.go | 899 +++++ third_party/quic-go/qlog/frame.go | 481 +++ third_party/quic-go/qlog/frame_test.go | 421 ++ third_party/quic-go/qlog/json_helper_test.go | 42 + third_party/quic-go/qlog/packet_header.go | 96 + .../quic-go/qlog/packet_header_test.go | 134 + third_party/quic-go/qlog/qlog_dir.go | 61 + third_party/quic-go/qlog/qlog_dir_test.go | 70 + third_party/quic-go/qlog/types.go | 305 ++ third_party/quic-go/qlog/types_test.go | 20 + .../quic-go/qlogwriter/jsontext/encoder.go | 324 ++ .../qlogwriter/jsontext/encoder_test.go | 405 ++ third_party/quic-go/qlogwriter/trace.go | 124 + third_party/quic-go/qlogwriter/trace_test.go | 115 + third_party/quic-go/qlogwriter/writer.go | 229 ++ third_party/quic-go/qlogwriter/writer_test.go | 116 + third_party/quic-go/quic_linux_test.go | 12 + third_party/quic-go/quic_test.go | 87 + third_party/quic-go/quicvarint/io.go | 98 + third_party/quic-go/quicvarint/io_test.go | 162 + third_party/quic-go/quicvarint/varint.go | 180 + third_party/quic-go/quicvarint/varint_test.go | 350 ++ third_party/quic-go/receive_stream.go | 578 +++ third_party/quic-go/receive_stream_test.go | 1140 ++++++ third_party/quic-go/retransmission_queue.go | 158 + .../quic-go/retransmission_queue_test.go | 131 + third_party/quic-go/send_conn.go | 136 + third_party/quic-go/send_conn_test.go | 135 + third_party/quic-go/send_queue.go | 114 + third_party/quic-go/send_queue_test.go | 207 + third_party/quic-go/send_stream.go | 915 +++++ third_party/quic-go/send_stream_test.go | 1875 +++++++++ third_party/quic-go/server.go | 1128 ++++++ third_party/quic-go/server_test.go | 1420 +++++++ third_party/quic-go/sni.go | 136 + third_party/quic-go/sni_test.go | 301 ++ third_party/quic-go/stateless_reset.go | 42 + third_party/quic-go/stateless_reset_test.go | 42 + third_party/quic-go/stream.go | 253 ++ third_party/quic-go/stream_test.go | 109 + third_party/quic-go/streams_map.go | 356 ++ third_party/quic-go/streams_map_incoming.go | 209 + .../quic-go/streams_map_incoming_test.go | 364 ++ third_party/quic-go/streams_map_outgoing.go | 253 ++ .../quic-go/streams_map_outgoing_test.go | 621 +++ third_party/quic-go/streams_map_test.go | 674 ++++ third_party/quic-go/sys_conn.go | 122 + third_party/quic-go/sys_conn_buffers.go | 68 + third_party/quic-go/sys_conn_buffers_write.go | 70 + third_party/quic-go/sys_conn_df.go | 22 + third_party/quic-go/sys_conn_df_darwin.go | 90 + .../quic-go/sys_conn_df_darwin_test.go | 101 + third_party/quic-go/sys_conn_df_linux.go | 42 + third_party/quic-go/sys_conn_df_windows.go | 52 + third_party/quic-go/sys_conn_helper_darwin.go | 38 + .../quic-go/sys_conn_helper_freebsd.go | 33 + third_party/quic-go/sys_conn_helper_linux.go | 156 + .../quic-go/sys_conn_helper_linux_test.go | 79 + .../quic-go/sys_conn_helper_nonlinux.go | 10 + .../quic-go/sys_conn_helper_nonlinux_test.go | 10 + third_party/quic-go/sys_conn_no_oob.go | 21 + third_party/quic-go/sys_conn_oob.go | 344 ++ third_party/quic-go/sys_conn_oob_test.go | 334 ++ third_party/quic-go/sys_conn_test.go | 32 + third_party/quic-go/sys_conn_windows.go | 42 + third_party/quic-go/sys_conn_windows_test.go | 30 + .../testutils/events/event_recorder.go | 94 + .../testutils/events/event_recorder_test.go | 101 + third_party/quic-go/testutils/frames.go | 26 + .../quic-go/testutils/simnet/README.md | 14 + third_party/quic-go/testutils/simnet/queue.go | 134 + .../quic-go/testutils/simnet/queue_test.go | 83 + .../quic-go/testutils/simnet/router.go | 149 + .../quic-go/testutils/simnet/simconn.go | 250 ++ .../quic-go/testutils/simnet/simconn_test.go | 187 + .../quic-go/testutils/simnet/simlink.go | 145 + .../quic-go/testutils/simnet/simlink_test.go | 153 + .../quic-go/testutils/simnet/simnet.go | 71 + .../testutils/simnet/simnet_synctest_test.go | 59 + third_party/quic-go/testutils/testutils.go | 108 + third_party/quic-go/token_store.go | 116 + third_party/quic-go/token_store_test.go | 79 + third_party/quic-go/transport.go | 866 ++++ third_party/quic-go/transport_test.go | 739 ++++ tools/notices/main.go | 320 ++ tools/notices/main_test.go | 87 + tools/vulnfilter/main.go | 140 + tools/vulnfilter/main_test.go | 43 + 606 files changed, 124623 insertions(+), 593 deletions(-) create mode 100644 THIRD_PARTY_NOTICES.md create mode 100644 docs/ACCELERATION.md create mode 100644 docs/assets/autocar-logo.svg create mode 100644 internal/hy2/auto.go create mode 100644 internal/hy2/client.go create mode 100644 internal/hy2/hy2_test.go create mode 100644 internal/hy2/server.go create mode 100644 internal/hy2/tls_policy.go create mode 100644 internal/proxy/socks5_udp_test.go create mode 100644 internal/tunnel/tls_limits_test.go create mode 100755 scripts/check-fork-provenance.sh create mode 100755 scripts/govulncheck.sh create mode 100644 third_party/hysteria-core/AUTOCAR_PATCHES.md create mode 100644 third_party/hysteria-core/LICENSE.md create mode 100644 third_party/hysteria-core/client/.mockery.yaml create mode 100644 third_party/hysteria-core/client/client.go create mode 100644 third_party/hysteria-core/client/config.go create mode 100644 third_party/hysteria-core/client/fast_open_test.go create mode 100644 third_party/hysteria-core/client/mock_udpIO.go create mode 100644 third_party/hysteria-core/client/reconnect.go create mode 100644 third_party/hysteria-core/client/udp.go create mode 100644 third_party/hysteria-core/client/udp_test.go create mode 100644 third_party/hysteria-core/errors/errors.go create mode 100644 third_party/hysteria-core/go.mod create mode 100644 third_party/hysteria-core/go.sum create mode 100644 third_party/hysteria-core/internal/congestion/bbr/bandwidth.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/bandwidth_sampler.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/bbr_sender.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/bbr_sender_test.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/clock.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/packet_number_indexed_queue.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/ringbuffer.go create mode 100644 third_party/hysteria-core/internal/congestion/bbr/windowed_filter.go create mode 100644 third_party/hysteria-core/internal/congestion/brutal/brutal.go create mode 100644 third_party/hysteria-core/internal/congestion/brutal/brutal_test.go create mode 100644 third_party/hysteria-core/internal/congestion/common/pacer.go create mode 100644 third_party/hysteria-core/internal/congestion/utils.go create mode 100644 third_party/hysteria-core/internal/frag/frag.go create mode 100644 third_party/hysteria-core/internal/frag/frag_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/.mockery.yaml create mode 100644 third_party/hysteria-core/internal/integration_tests/chrome_parrot_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/close_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/hook_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/masq_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_Authenticator.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_Conn.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_EventLogger.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_Outbound.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_RequestHook.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_TrafficLogger.go create mode 100644 third_party/hysteria-core/internal/integration_tests/mocks/mock_UDPConn.go create mode 100644 third_party/hysteria-core/internal/integration_tests/smoke_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/stress_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/trafficlogger_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/udp_acl_test.go create mode 100644 third_party/hysteria-core/internal/integration_tests/utils_test.go create mode 100644 third_party/hysteria-core/internal/pmtud/avail.go create mode 100644 third_party/hysteria-core/internal/pmtud/unavail.go create mode 100644 third_party/hysteria-core/internal/protocol/http.go create mode 100644 third_party/hysteria-core/internal/protocol/padding.go create mode 100644 third_party/hysteria-core/internal/protocol/proxy.go create mode 100644 third_party/hysteria-core/internal/protocol/proxy_test.go create mode 100644 third_party/hysteria-core/internal/utils/atomic.go create mode 100644 third_party/hysteria-core/internal/utils/qstream.go create mode 100644 third_party/hysteria-core/server/.mockery.yaml create mode 100644 third_party/hysteria-core/server/config.go create mode 100644 third_party/hysteria-core/server/copy.go create mode 100644 third_party/hysteria-core/server/copy_benchmark_test.go create mode 100644 third_party/hysteria-core/server/mock_UDPConn.go create mode 100644 third_party/hysteria-core/server/mock_udpEventLogger.go create mode 100644 third_party/hysteria-core/server/mock_udpIO.go create mode 100644 third_party/hysteria-core/server/resource_limits_test.go create mode 100644 third_party/hysteria-core/server/server.go create mode 100644 third_party/hysteria-core/server/udp.go create mode 100644 third_party/hysteria-core/server/udp_test.go create mode 100644 third_party/quic-go/.clusterfuzzlite/Dockerfile create mode 100644 third_party/quic-go/.clusterfuzzlite/build.sh create mode 100644 third_party/quic-go/.clusterfuzzlite/project.yaml create mode 100644 third_party/quic-go/.githooks/README.md create mode 100644 third_party/quic-go/.githooks/pre-commit create mode 100644 third_party/quic-go/.github/FUNDING.yml create mode 100644 third_party/quic-go/.github/dependabot.yml create mode 100644 third_party/quic-go/.github/workflows/build-interop-docker.yml create mode 100644 third_party/quic-go/.github/workflows/clusterfuzz-coverage.yml create mode 100644 third_party/quic-go/.github/workflows/clusterfuzz-lite-batch.yml create mode 100644 third_party/quic-go/.github/workflows/clusterfuzz-lite-pr.yml create mode 100644 third_party/quic-go/.github/workflows/clusterfuzz-lite-prune.yml create mode 100644 third_party/quic-go/.github/workflows/codspeed.yml create mode 100644 third_party/quic-go/.github/workflows/cross-compile.sh create mode 100644 third_party/quic-go/.github/workflows/cross-compile.yml create mode 100644 third_party/quic-go/.github/workflows/go-generate.sh create mode 100644 third_party/quic-go/.github/workflows/govulncheck.yml create mode 100644 third_party/quic-go/.github/workflows/integration.yml create mode 100644 third_party/quic-go/.github/workflows/lint.yml create mode 100644 third_party/quic-go/.github/workflows/unit.yml create mode 100644 third_party/quic-go/.gitignore create mode 100644 third_party/quic-go/.golangci.yml create mode 100644 third_party/quic-go/AUTOCAR_PATCHES.md create mode 100644 third_party/quic-go/FIPS140.md create mode 100644 third_party/quic-go/FUZZING.md create mode 100644 third_party/quic-go/LICENSE create mode 100644 third_party/quic-go/README.md create mode 100644 third_party/quic-go/SECURITY.md create mode 100644 third_party/quic-go/assets/LICENSE.md create mode 100644 third_party/quic-go/assets/logo.svg create mode 100644 third_party/quic-go/assets/quic-go-logo.png create mode 100644 third_party/quic-go/buffer_pool.go create mode 100644 third_party/quic-go/buffer_pool_test.go create mode 100644 third_party/quic-go/chrome_parrot.go create mode 100644 third_party/quic-go/client.go create mode 100644 third_party/quic-go/client_test.go create mode 100644 third_party/quic-go/closed_conn.go create mode 100644 third_party/quic-go/closed_conn_test.go create mode 100644 third_party/quic-go/codecov.yml create mode 100644 third_party/quic-go/config.go create mode 100644 third_party/quic-go/config_test.go create mode 100644 third_party/quic-go/congestion/interface.go create mode 100644 third_party/quic-go/conn_id_generator.go create mode 100644 third_party/quic-go/conn_id_generator_test.go create mode 100644 third_party/quic-go/conn_id_manager.go create mode 100644 third_party/quic-go/conn_id_manager_test.go create mode 100644 third_party/quic-go/conn_wrapped_test.go create mode 100644 third_party/quic-go/connection.go create mode 100644 third_party/quic-go/connection_logging.go create mode 100644 third_party/quic-go/connection_logging_test.go create mode 100644 third_party/quic-go/connection_test.go create mode 100644 third_party/quic-go/crypto_stream.go create mode 100644 third_party/quic-go/crypto_stream_manager.go create mode 100644 third_party/quic-go/crypto_stream_manager_test.go create mode 100644 third_party/quic-go/crypto_stream_test.go create mode 100644 third_party/quic-go/datagram_queue.go create mode 100644 third_party/quic-go/datagram_queue_test.go create mode 100644 third_party/quic-go/errors.go create mode 100644 third_party/quic-go/errors_test.go create mode 100644 third_party/quic-go/example/client/main.go create mode 100644 third_party/quic-go/example/echo/echo.go create mode 100644 third_party/quic-go/example/main.go create mode 100644 third_party/quic-go/flow_controller_base.go create mode 100644 third_party/quic-go/flow_controller_connection.go create mode 100644 third_party/quic-go/flow_controller_connection_test.go create mode 100644 third_party/quic-go/flow_controller_stream.go create mode 100644 third_party/quic-go/flow_controller_stream_test.go create mode 100644 third_party/quic-go/flow_controller_test_helpers_test.go create mode 100644 third_party/quic-go/frame_sorter.go create mode 100644 third_party/quic-go/frame_sorter_test.go create mode 100644 third_party/quic-go/framer.go create mode 100644 third_party/quic-go/framer_test.go create mode 100644 third_party/quic-go/go.mod create mode 100644 third_party/quic-go/go.sum create mode 100644 third_party/quic-go/http3/README.md create mode 100644 third_party/quic-go/http3/body.go create mode 100644 third_party/quic-go/http3/body_test.go create mode 100644 third_party/quic-go/http3/capsule.go create mode 100644 third_party/quic-go/http3/capsule_test.go create mode 100644 third_party/quic-go/http3/client.go create mode 100644 third_party/quic-go/http3/client_test.go create mode 100644 third_party/quic-go/http3/conn.go create mode 100644 third_party/quic-go/http3/conn_test.go create mode 100644 third_party/quic-go/http3/error.go create mode 100644 third_party/quic-go/http3/error_codes.go create mode 100644 third_party/quic-go/http3/error_codes_test.go create mode 100644 third_party/quic-go/http3/error_test.go create mode 100644 third_party/quic-go/http3/frames.go create mode 100644 third_party/quic-go/http3/frames_test.go create mode 100644 third_party/quic-go/http3/gzip_reader.go create mode 100644 third_party/quic-go/http3/headers.go create mode 100644 third_party/quic-go/http3/headers_test.go create mode 100644 third_party/quic-go/http3/http3_helper_test.go create mode 100644 third_party/quic-go/http3/internal/testdata/cert.go create mode 100644 third_party/quic-go/http3/internal/testdata/cert_test.go create mode 100644 third_party/quic-go/http3/ip_addr.go create mode 100644 third_party/quic-go/http3/mock_clientconn_test.go create mode 100644 third_party/quic-go/http3/mock_datagram_stream_test.go create mode 100644 third_party/quic-go/http3/mock_quic_listener_test.go create mode 100644 third_party/quic-go/http3/mockgen.go create mode 100644 third_party/quic-go/http3/qlog.go create mode 100644 third_party/quic-go/http3/qlog/event.go create mode 100644 third_party/quic-go/http3/qlog/event_test.go create mode 100644 third_party/quic-go/http3/qlog/frame.go create mode 100644 third_party/quic-go/http3/qlog/frame_test.go create mode 100644 third_party/quic-go/http3/qlog/qlog_dir.go create mode 100644 third_party/quic-go/http3/qlog/qlog_dir_test.go create mode 100644 third_party/quic-go/http3/request_writer.go create mode 100644 third_party/quic-go/http3/request_writer_test.go create mode 100644 third_party/quic-go/http3/response_writer.go create mode 100644 third_party/quic-go/http3/response_writer_test.go create mode 100644 third_party/quic-go/http3/server.go create mode 100644 third_party/quic-go/http3/server_conn.go create mode 100644 third_party/quic-go/http3/server_test.go create mode 100644 third_party/quic-go/http3/state_tracking_stream.go create mode 100644 third_party/quic-go/http3/state_tracking_stream_test.go create mode 100644 third_party/quic-go/http3/stream.go create mode 100644 third_party/quic-go/http3/stream_test.go create mode 100644 third_party/quic-go/http3/trace.go create mode 100644 third_party/quic-go/http3/transport.go create mode 100644 third_party/quic-go/http3/transport_test.go create mode 100644 third_party/quic-go/integrationtests/self/benchmark_test.go create mode 100644 third_party/quic-go/integrationtests/self/cancelation_test.go create mode 100644 third_party/quic-go/integrationtests/self/chrome_parrot_test.go create mode 100644 third_party/quic-go/integrationtests/self/close_test.go create mode 100644 third_party/quic-go/integrationtests/self/conn_id_test.go create mode 100644 third_party/quic-go/integrationtests/self/connection_migration_test.go create mode 100644 third_party/quic-go/integrationtests/self/datagram_test.go create mode 100644 third_party/quic-go/integrationtests/self/deadline_test.go create mode 100644 third_party/quic-go/integrationtests/self/drop_test.go create mode 100644 third_party/quic-go/integrationtests/self/early_data_test.go create mode 100644 third_party/quic-go/integrationtests/self/handshake_context_test.go create mode 100644 third_party/quic-go/integrationtests/self/handshake_drop_test.go create mode 100644 third_party/quic-go/integrationtests/self/handshake_rtt_test.go create mode 100644 third_party/quic-go/integrationtests/self/handshake_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_datagram_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_hotswap_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_qlog_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_raw_conn_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_shutdown_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_test.go create mode 100644 third_party/quic-go/integrationtests/self/http_trace_test.go create mode 100644 third_party/quic-go/integrationtests/self/key_update_test.go create mode 100644 third_party/quic-go/integrationtests/self/mitm_test.go create mode 100644 third_party/quic-go/integrationtests/self/mtu_test.go create mode 100644 third_party/quic-go/integrationtests/self/multiplex_test.go create mode 100644 third_party/quic-go/integrationtests/self/nat_rebinding_test.go create mode 100644 third_party/quic-go/integrationtests/self/packetization_test.go create mode 100644 third_party/quic-go/integrationtests/self/qlog_dir_test.go create mode 100644 third_party/quic-go/integrationtests/self/qlog_test.go create mode 100644 third_party/quic-go/integrationtests/self/resumption_test.go create mode 100644 third_party/quic-go/integrationtests/self/rtt_test.go create mode 100644 third_party/quic-go/integrationtests/self/self_go124_test.go create mode 100644 third_party/quic-go/integrationtests/self/self_go125_test.go create mode 100644 third_party/quic-go/integrationtests/self/self_suite_linux_test.go create mode 100644 third_party/quic-go/integrationtests/self/self_suite_others_test.go create mode 100644 third_party/quic-go/integrationtests/self/self_test.go create mode 100644 third_party/quic-go/integrationtests/self/simnet_helper_test.go create mode 100644 third_party/quic-go/integrationtests/self/stateless_reset_test.go create mode 100644 third_party/quic-go/integrationtests/self/stream_test.go create mode 100644 third_party/quic-go/integrationtests/self/timeout_test.go create mode 100644 third_party/quic-go/integrationtests/self/zero_rtt_test.go create mode 100644 third_party/quic-go/integrationtests/tools/crypto.go create mode 100644 third_party/quic-go/integrationtests/tools/crypto_test.go create mode 100644 third_party/quic-go/integrationtests/tools/israce/norace.go create mode 100644 third_party/quic-go/integrationtests/tools/israce/race.go create mode 100644 third_party/quic-go/integrationtests/tools/proxy/proxy.go create mode 100644 third_party/quic-go/integrationtests/tools/proxy/proxy_test.go create mode 100644 third_party/quic-go/integrationtests/tools/qlog.go create mode 100644 third_party/quic-go/integrationtests/versionnegotiation/handshake_test.go create mode 100644 third_party/quic-go/integrationtests/versionnegotiation/rtt_test.go create mode 100644 third_party/quic-go/integrationtests/versionnegotiation/test_helper_test.go create mode 100644 third_party/quic-go/interface.go create mode 100644 third_party/quic-go/internal/ackhandler/ack_eliciting.go create mode 100644 third_party/quic-go/internal/ackhandler/ack_eliciting_test.go create mode 100644 third_party/quic-go/internal/ackhandler/cc_adapter.go create mode 100644 third_party/quic-go/internal/ackhandler/cc_adapter_ex.go create mode 100644 third_party/quic-go/internal/ackhandler/ecn.go create mode 100644 third_party/quic-go/internal/ackhandler/ecn_test.go create mode 100644 third_party/quic-go/internal/ackhandler/frame.go create mode 100644 third_party/quic-go/internal/ackhandler/interfaces.go create mode 100644 third_party/quic-go/internal/ackhandler/lost_packet_tracker.go create mode 100644 third_party/quic-go/internal/ackhandler/lost_packet_tracker_test.go create mode 100644 third_party/quic-go/internal/ackhandler/mock_ecn_handler_test.go create mode 100644 third_party/quic-go/internal/ackhandler/mockgen.go create mode 100644 third_party/quic-go/internal/ackhandler/packet.go create mode 100644 third_party/quic-go/internal/ackhandler/packet_number_generator.go create mode 100644 third_party/quic-go/internal/ackhandler/packet_number_generator_test.go create mode 100644 third_party/quic-go/internal/ackhandler/received_packet_handler.go create mode 100644 third_party/quic-go/internal/ackhandler/received_packet_handler_test.go create mode 100644 third_party/quic-go/internal/ackhandler/received_packet_history.go create mode 100644 third_party/quic-go/internal/ackhandler/received_packet_history_test.go create mode 100644 third_party/quic-go/internal/ackhandler/received_packet_tracker.go create mode 100644 third_party/quic-go/internal/ackhandler/received_packet_tracker_test.go create mode 100644 third_party/quic-go/internal/ackhandler/send_mode.go create mode 100644 third_party/quic-go/internal/ackhandler/send_mode_test.go create mode 100644 third_party/quic-go/internal/ackhandler/sent_packet_handler.go create mode 100644 third_party/quic-go/internal/ackhandler/sent_packet_handler_test.go create mode 100644 third_party/quic-go/internal/ackhandler/sent_packet_history.go create mode 100644 third_party/quic-go/internal/ackhandler/sent_packet_history_test.go create mode 100644 third_party/quic-go/internal/congestion/bandwidth.go create mode 100644 third_party/quic-go/internal/congestion/bandwidth_test.go create mode 100644 third_party/quic-go/internal/congestion/clock.go create mode 100644 third_party/quic-go/internal/congestion/cubic.go create mode 100644 third_party/quic-go/internal/congestion/cubic_sender.go create mode 100644 third_party/quic-go/internal/congestion/cubic_sender_test.go create mode 100644 third_party/quic-go/internal/congestion/cubic_test.go create mode 100644 third_party/quic-go/internal/congestion/hybrid_slow_start.go create mode 100644 third_party/quic-go/internal/congestion/hybrid_slow_start_test.go create mode 100644 third_party/quic-go/internal/congestion/interface.go create mode 100644 third_party/quic-go/internal/congestion/pacer.go create mode 100644 third_party/quic-go/internal/congestion/pacer_test.go create mode 100644 third_party/quic-go/internal/handshake/aead.go create mode 100644 third_party/quic-go/internal/handshake/aead_test.go create mode 100644 third_party/quic-go/internal/handshake/chrome_client_hello.go create mode 100644 third_party/quic-go/internal/handshake/cipher_suite.go create mode 100644 third_party/quic-go/internal/handshake/cipher_suite_fips140.go create mode 100644 third_party/quic-go/internal/handshake/crypto_setup.go create mode 100644 third_party/quic-go/internal/handshake/crypto_setup_test.go create mode 100644 third_party/quic-go/internal/handshake/fake_conn.go create mode 100644 third_party/quic-go/internal/handshake/fips140_go126.go create mode 100644 third_party/quic-go/internal/handshake/fips140_legacy.go create mode 100644 third_party/quic-go/internal/handshake/handshake_fuzz_test.go create mode 100644 third_party/quic-go/internal/handshake/handshake_helpers_test.go create mode 100644 third_party/quic-go/internal/handshake/header_protector.go create mode 100644 third_party/quic-go/internal/handshake/hkdf.go create mode 100644 third_party/quic-go/internal/handshake/hkdf_test.go create mode 100644 third_party/quic-go/internal/handshake/initial_aead.go create mode 100644 third_party/quic-go/internal/handshake/initial_aead_test.go create mode 100644 third_party/quic-go/internal/handshake/interface.go create mode 100644 third_party/quic-go/internal/handshake/quic_event_go125.go create mode 100644 third_party/quic-go/internal/handshake/quic_event_go126.go create mode 100644 third_party/quic-go/internal/handshake/retry_go125.go create mode 100644 third_party/quic-go/internal/handshake/retry_go126.go create mode 100644 third_party/quic-go/internal/handshake/retry_test.go create mode 100644 third_party/quic-go/internal/handshake/session_ticket.go create mode 100644 third_party/quic-go/internal/handshake/session_ticket_test.go create mode 100644 third_party/quic-go/internal/handshake/tls_config_go126.go create mode 100644 third_party/quic-go/internal/handshake/tls_config_go126_test.go create mode 100644 third_party/quic-go/internal/handshake/tls_config_go127.go create mode 100644 third_party/quic-go/internal/handshake/tls_conn.go create mode 100644 third_party/quic-go/internal/handshake/tls_conn_utls.go create mode 100644 third_party/quic-go/internal/handshake/token_generator.go create mode 100644 third_party/quic-go/internal/handshake/token_generator_test.go create mode 100644 third_party/quic-go/internal/handshake/token_protector.go create mode 100644 third_party/quic-go/internal/handshake/token_protector_test.go create mode 100644 third_party/quic-go/internal/handshake/updatable_aead.go create mode 100644 third_party/quic-go/internal/handshake/updatable_aead_test.go create mode 100644 third_party/quic-go/internal/mocks/ackhandler/sent_packet_handler.go create mode 100644 third_party/quic-go/internal/mocks/congestion.go create mode 100644 third_party/quic-go/internal/mocks/crypto_setup.go create mode 100644 third_party/quic-go/internal/mocks/long_header_opener.go create mode 100644 third_party/quic-go/internal/mocks/mockgen.go create mode 100644 third_party/quic-go/internal/mocks/short_header_opener.go create mode 100644 third_party/quic-go/internal/mocks/short_header_sealer.go create mode 100644 third_party/quic-go/internal/monotime/time.go create mode 100644 third_party/quic-go/internal/monotime/time_test.go create mode 100644 third_party/quic-go/internal/protocol/connection_id.go create mode 100644 third_party/quic-go/internal/protocol/connection_id_test.go create mode 100644 third_party/quic-go/internal/protocol/encryption_level.go create mode 100644 third_party/quic-go/internal/protocol/encryption_level_test.go create mode 100644 third_party/quic-go/internal/protocol/key_phase.go create mode 100644 third_party/quic-go/internal/protocol/key_phase_test.go create mode 100644 third_party/quic-go/internal/protocol/packet_number.go create mode 100644 third_party/quic-go/internal/protocol/packet_number_test.go create mode 100644 third_party/quic-go/internal/protocol/params.go create mode 100644 third_party/quic-go/internal/protocol/params_test.go create mode 100644 third_party/quic-go/internal/protocol/perspective.go create mode 100644 third_party/quic-go/internal/protocol/perspective_test.go create mode 100644 third_party/quic-go/internal/protocol/protocol.go create mode 100644 third_party/quic-go/internal/protocol/protocol_test.go create mode 100644 third_party/quic-go/internal/protocol/stream.go create mode 100644 third_party/quic-go/internal/protocol/stream_test.go create mode 100644 third_party/quic-go/internal/protocol/version.go create mode 100644 third_party/quic-go/internal/protocol/version_test.go create mode 100644 third_party/quic-go/internal/qerr/error_codes.go create mode 100644 third_party/quic-go/internal/qerr/errorcodes_test.go create mode 100644 third_party/quic-go/internal/qerr/errors.go create mode 100644 third_party/quic-go/internal/qerr/errors_test.go create mode 100644 third_party/quic-go/internal/qtls/cipher_suite.go create mode 100644 third_party/quic-go/internal/qtls/cipher_suite_test.go create mode 100644 third_party/quic-go/internal/testdata/cert.go create mode 100644 third_party/quic-go/internal/testdata/cert_test.go create mode 100644 third_party/quic-go/internal/utils/buffered_write_closer.go create mode 100644 third_party/quic-go/internal/utils/buffered_write_closer_test.go create mode 100644 third_party/quic-go/internal/utils/connstats.go create mode 100644 third_party/quic-go/internal/utils/linkedlist/README.md create mode 100644 third_party/quic-go/internal/utils/linkedlist/linkedlist.go create mode 100644 third_party/quic-go/internal/utils/log.go create mode 100644 third_party/quic-go/internal/utils/log_test.go create mode 100644 third_party/quic-go/internal/utils/rand.go create mode 100644 third_party/quic-go/internal/utils/rand_test.go create mode 100644 third_party/quic-go/internal/utils/ringbuffer/ringbuffer.go create mode 100644 third_party/quic-go/internal/utils/ringbuffer/ringbuffer_bench_test.go create mode 100644 third_party/quic-go/internal/utils/ringbuffer/ringbuffer_test.go create mode 100644 third_party/quic-go/internal/utils/rtt_stats.go create mode 100644 third_party/quic-go/internal/utils/rtt_stats_test.go create mode 100644 third_party/quic-go/internal/utils/streamframe_interval.go create mode 100644 third_party/quic-go/internal/utils/tree/tree.go create mode 100644 third_party/quic-go/internal/utils/tree/tree_match_test.go create mode 100644 third_party/quic-go/internal/utils/tree/tree_test.go create mode 100644 third_party/quic-go/internal/wire/ack_frame.go create mode 100644 third_party/quic-go/internal/wire/ack_frame_test.go create mode 100644 third_party/quic-go/internal/wire/ack_frequency_frame.go create mode 100644 third_party/quic-go/internal/wire/ack_frequency_frame_test.go create mode 100644 third_party/quic-go/internal/wire/ack_range.go create mode 100644 third_party/quic-go/internal/wire/ack_range_test.go create mode 100644 third_party/quic-go/internal/wire/connection_close_frame.go create mode 100644 third_party/quic-go/internal/wire/connection_close_frame_test.go create mode 100644 third_party/quic-go/internal/wire/crypto_frame.go create mode 100644 third_party/quic-go/internal/wire/crypto_frame_test.go create mode 100644 third_party/quic-go/internal/wire/data_blocked_frame.go create mode 100644 third_party/quic-go/internal/wire/data_blocked_frame_test.go create mode 100644 third_party/quic-go/internal/wire/datagram_frame.go create mode 100644 third_party/quic-go/internal/wire/datagram_frame_test.go create mode 100644 third_party/quic-go/internal/wire/extended_header.go create mode 100644 third_party/quic-go/internal/wire/extended_header_test.go create mode 100644 third_party/quic-go/internal/wire/frame.go create mode 100644 third_party/quic-go/internal/wire/frame_parser.go create mode 100644 third_party/quic-go/internal/wire/frame_parser_test.go create mode 100644 third_party/quic-go/internal/wire/frame_test.go create mode 100644 third_party/quic-go/internal/wire/frame_type.go create mode 100644 third_party/quic-go/internal/wire/frame_type_test.go create mode 100644 third_party/quic-go/internal/wire/handshake_done_frame.go create mode 100644 third_party/quic-go/internal/wire/handshake_done_frame_test.go create mode 100644 third_party/quic-go/internal/wire/header.go create mode 100644 third_party/quic-go/internal/wire/header_test.go create mode 100644 third_party/quic-go/internal/wire/immediate_ack_frame.go create mode 100644 third_party/quic-go/internal/wire/immediate_ack_frame_test.go create mode 100644 third_party/quic-go/internal/wire/log.go create mode 100644 third_party/quic-go/internal/wire/log_test.go create mode 100644 third_party/quic-go/internal/wire/max_data_frame.go create mode 100644 third_party/quic-go/internal/wire/max_data_frame_test.go create mode 100644 third_party/quic-go/internal/wire/max_stream_data_frame.go create mode 100644 third_party/quic-go/internal/wire/max_stream_data_frame_test.go create mode 100644 third_party/quic-go/internal/wire/max_streams_frame.go create mode 100644 third_party/quic-go/internal/wire/max_streams_frame_test.go create mode 100644 third_party/quic-go/internal/wire/new_connection_id_frame.go create mode 100644 third_party/quic-go/internal/wire/new_connection_id_frame_test.go create mode 100644 third_party/quic-go/internal/wire/new_token_frame.go create mode 100644 third_party/quic-go/internal/wire/new_token_frame_test.go create mode 100644 third_party/quic-go/internal/wire/path_challenge_frame.go create mode 100644 third_party/quic-go/internal/wire/path_challenge_frame_test.go create mode 100644 third_party/quic-go/internal/wire/path_response_frame.go create mode 100644 third_party/quic-go/internal/wire/path_response_frame_test.go create mode 100644 third_party/quic-go/internal/wire/ping_frame.go create mode 100644 third_party/quic-go/internal/wire/ping_frame_test.go create mode 100644 third_party/quic-go/internal/wire/pool.go create mode 100644 third_party/quic-go/internal/wire/pool_test.go create mode 100644 third_party/quic-go/internal/wire/reset_stream_frame.go create mode 100644 third_party/quic-go/internal/wire/reset_stream_frame_test.go create mode 100644 third_party/quic-go/internal/wire/retire_connection_id_frame.go create mode 100644 third_party/quic-go/internal/wire/retire_connection_id_frame_test.go create mode 100644 third_party/quic-go/internal/wire/short_header.go create mode 100644 third_party/quic-go/internal/wire/short_header_test.go create mode 100644 third_party/quic-go/internal/wire/stop_sending_frame.go create mode 100644 third_party/quic-go/internal/wire/stop_sending_frame_test.go create mode 100644 third_party/quic-go/internal/wire/stream_data_blocked_frame.go create mode 100644 third_party/quic-go/internal/wire/stream_data_blocked_frame_test.go create mode 100644 third_party/quic-go/internal/wire/stream_frame.go create mode 100644 third_party/quic-go/internal/wire/stream_frame_test.go create mode 100644 third_party/quic-go/internal/wire/streams_blocked_frame.go create mode 100644 third_party/quic-go/internal/wire/streams_blocked_frame_test.go create mode 100644 third_party/quic-go/internal/wire/test_helpers_test.go create mode 100644 third_party/quic-go/internal/wire/transport_parameter_test.go create mode 100644 third_party/quic-go/internal/wire/transport_parameters.go create mode 100644 third_party/quic-go/internal/wire/transport_parameters_chrome.go create mode 100644 third_party/quic-go/internal/wire/transport_parameters_chrome_test.go create mode 100644 third_party/quic-go/internal/wire/version_negotiation.go create mode 100644 third_party/quic-go/internal/wire/version_negotiation_test.go create mode 100644 third_party/quic-go/interop/Dockerfile create mode 100644 third_party/quic-go/interop/client/main.go create mode 100644 third_party/quic-go/interop/http09/client.go create mode 100644 third_party/quic-go/interop/http09/http_test.go create mode 100644 third_party/quic-go/interop/http09/server.go create mode 100644 third_party/quic-go/interop/run_endpoint.sh create mode 100644 third_party/quic-go/interop/server/main.go create mode 100644 third_party/quic-go/interop/utils/logging.go create mode 100644 third_party/quic-go/metrics/dashboards/README.md create mode 100644 third_party/quic-go/metrics/dashboards/datasources.yml create mode 100644 third_party/quic-go/metrics/dashboards/docker-compose.yml create mode 100644 third_party/quic-go/metrics/dashboards/prometheus.yml create mode 100644 third_party/quic-go/metrics/dashboards/quic-go.json create mode 100644 third_party/quic-go/mock_ack_frame_source_test.go create mode 100644 third_party/quic-go/mock_conn_runner_test.go create mode 100644 third_party/quic-go/mock_frame_source_test.go create mode 100644 third_party/quic-go/mock_mtu_discoverer_test.go create mode 100644 third_party/quic-go/mock_packer_test.go create mode 100644 third_party/quic-go/mock_packet_handler_test.go create mode 100644 third_party/quic-go/mock_packetconn_test.go create mode 100644 third_party/quic-go/mock_raw_conn_test.go create mode 100644 third_party/quic-go/mock_sealing_manager_test.go create mode 100644 third_party/quic-go/mock_send_conn_test.go create mode 100644 third_party/quic-go/mock_sender_test.go create mode 100644 third_party/quic-go/mock_stream_control_frame_getter_test.go create mode 100644 third_party/quic-go/mock_stream_frame_getter_test.go create mode 100644 third_party/quic-go/mock_stream_sender_test.go create mode 100644 third_party/quic-go/mock_unpacker_test.go create mode 100644 third_party/quic-go/mockgen.go create mode 100644 third_party/quic-go/module_rename.sh create mode 100644 third_party/quic-go/monotime/time.go create mode 100644 third_party/quic-go/mtu_discoverer.go create mode 100644 third_party/quic-go/mtu_discoverer_test.go create mode 100644 third_party/quic-go/oss-fuzz.sh create mode 100644 third_party/quic-go/packet_packer.go create mode 100644 third_party/quic-go/packet_packer_chaos.go create mode 100644 third_party/quic-go/packet_packer_chaos_test.go create mode 100644 third_party/quic-go/packet_packer_test.go create mode 100644 third_party/quic-go/packet_unpacker.go create mode 100644 third_party/quic-go/packet_unpacker_test.go create mode 100644 third_party/quic-go/path_manager.go create mode 100644 third_party/quic-go/path_manager_outgoing.go create mode 100644 third_party/quic-go/path_manager_outgoing_test.go create mode 100644 third_party/quic-go/path_manager_test.go create mode 100644 third_party/quic-go/qlog/benchmark_test.go create mode 100644 third_party/quic-go/qlog/event.go create mode 100644 third_party/quic-go/qlog/event_test.go create mode 100644 third_party/quic-go/qlog/frame.go create mode 100644 third_party/quic-go/qlog/frame_test.go create mode 100644 third_party/quic-go/qlog/json_helper_test.go create mode 100644 third_party/quic-go/qlog/packet_header.go create mode 100644 third_party/quic-go/qlog/packet_header_test.go create mode 100644 third_party/quic-go/qlog/qlog_dir.go create mode 100644 third_party/quic-go/qlog/qlog_dir_test.go create mode 100644 third_party/quic-go/qlog/types.go create mode 100644 third_party/quic-go/qlog/types_test.go create mode 100644 third_party/quic-go/qlogwriter/jsontext/encoder.go create mode 100644 third_party/quic-go/qlogwriter/jsontext/encoder_test.go create mode 100644 third_party/quic-go/qlogwriter/trace.go create mode 100644 third_party/quic-go/qlogwriter/trace_test.go create mode 100644 third_party/quic-go/qlogwriter/writer.go create mode 100644 third_party/quic-go/qlogwriter/writer_test.go create mode 100644 third_party/quic-go/quic_linux_test.go create mode 100644 third_party/quic-go/quic_test.go create mode 100644 third_party/quic-go/quicvarint/io.go create mode 100644 third_party/quic-go/quicvarint/io_test.go create mode 100644 third_party/quic-go/quicvarint/varint.go create mode 100644 third_party/quic-go/quicvarint/varint_test.go create mode 100644 third_party/quic-go/receive_stream.go create mode 100644 third_party/quic-go/receive_stream_test.go create mode 100644 third_party/quic-go/retransmission_queue.go create mode 100644 third_party/quic-go/retransmission_queue_test.go create mode 100644 third_party/quic-go/send_conn.go create mode 100644 third_party/quic-go/send_conn_test.go create mode 100644 third_party/quic-go/send_queue.go create mode 100644 third_party/quic-go/send_queue_test.go create mode 100644 third_party/quic-go/send_stream.go create mode 100644 third_party/quic-go/send_stream_test.go create mode 100644 third_party/quic-go/server.go create mode 100644 third_party/quic-go/server_test.go create mode 100644 third_party/quic-go/sni.go create mode 100644 third_party/quic-go/sni_test.go create mode 100644 third_party/quic-go/stateless_reset.go create mode 100644 third_party/quic-go/stateless_reset_test.go create mode 100644 third_party/quic-go/stream.go create mode 100644 third_party/quic-go/stream_test.go create mode 100644 third_party/quic-go/streams_map.go create mode 100644 third_party/quic-go/streams_map_incoming.go create mode 100644 third_party/quic-go/streams_map_incoming_test.go create mode 100644 third_party/quic-go/streams_map_outgoing.go create mode 100644 third_party/quic-go/streams_map_outgoing_test.go create mode 100644 third_party/quic-go/streams_map_test.go create mode 100644 third_party/quic-go/sys_conn.go create mode 100644 third_party/quic-go/sys_conn_buffers.go create mode 100644 third_party/quic-go/sys_conn_buffers_write.go create mode 100644 third_party/quic-go/sys_conn_df.go create mode 100644 third_party/quic-go/sys_conn_df_darwin.go create mode 100644 third_party/quic-go/sys_conn_df_darwin_test.go create mode 100644 third_party/quic-go/sys_conn_df_linux.go create mode 100644 third_party/quic-go/sys_conn_df_windows.go create mode 100644 third_party/quic-go/sys_conn_helper_darwin.go create mode 100644 third_party/quic-go/sys_conn_helper_freebsd.go create mode 100644 third_party/quic-go/sys_conn_helper_linux.go create mode 100644 third_party/quic-go/sys_conn_helper_linux_test.go create mode 100644 third_party/quic-go/sys_conn_helper_nonlinux.go create mode 100644 third_party/quic-go/sys_conn_helper_nonlinux_test.go create mode 100644 third_party/quic-go/sys_conn_no_oob.go create mode 100644 third_party/quic-go/sys_conn_oob.go create mode 100644 third_party/quic-go/sys_conn_oob_test.go create mode 100644 third_party/quic-go/sys_conn_test.go create mode 100644 third_party/quic-go/sys_conn_windows.go create mode 100644 third_party/quic-go/sys_conn_windows_test.go create mode 100644 third_party/quic-go/testutils/events/event_recorder.go create mode 100644 third_party/quic-go/testutils/events/event_recorder_test.go create mode 100644 third_party/quic-go/testutils/frames.go create mode 100644 third_party/quic-go/testutils/simnet/README.md create mode 100644 third_party/quic-go/testutils/simnet/queue.go create mode 100644 third_party/quic-go/testutils/simnet/queue_test.go create mode 100644 third_party/quic-go/testutils/simnet/router.go create mode 100644 third_party/quic-go/testutils/simnet/simconn.go create mode 100644 third_party/quic-go/testutils/simnet/simconn_test.go create mode 100644 third_party/quic-go/testutils/simnet/simlink.go create mode 100644 third_party/quic-go/testutils/simnet/simlink_test.go create mode 100644 third_party/quic-go/testutils/simnet/simnet.go create mode 100644 third_party/quic-go/testutils/simnet/simnet_synctest_test.go create mode 100644 third_party/quic-go/testutils/testutils.go create mode 100644 third_party/quic-go/token_store.go create mode 100644 third_party/quic-go/token_store_test.go create mode 100644 third_party/quic-go/transport.go create mode 100644 third_party/quic-go/transport_test.go create mode 100644 tools/notices/main.go create mode 100644 tools/notices/main_test.go create mode 100644 tools/vulnfilter/main.go create mode 100644 tools/vulnfilter/main_test.go diff --git a/.github/dependabot.yml b/.github/dependabot.yml index b659d6c..f5ad1ae 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -13,6 +13,30 @@ updates: - minor - patch + - package-ecosystem: gomod + directory: /third_party/hysteria-core + schedule: + interval: weekly + day: monday + open-pull-requests-limit: 3 + groups: + hardened-core-minor-and-patch: + update-types: + - minor + - patch + + - package-ecosystem: gomod + directory: /third_party/quic-go + schedule: + interval: weekly + day: monday + open-pull-requests-limit: 3 + groups: + hardened-quic-minor-and-patch: + update-types: + - minor + - patch + - package-ecosystem: github-actions directory: / schedule: @@ -24,6 +48,11 @@ updates: update-types: - minor - patch + exclude-patterns: + - github/codeql-action/* + codeql-actions: + patterns: + - github/codeql-action/* - package-ecosystem: docker directory: / diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c09d545..670295e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -22,21 +22,36 @@ jobs: matrix: go: [1.25.x, 1.27.x] steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 with: go-version: ${{ matrix.go }} cache: true - name: Verify formatting run: make fmt-check + - name: Verify local-fork provenance + run: make fork-provenance-check - name: Verify module files run: make mod-check + - name: Verify third-party notices + run: make notices-check - name: Vet run: make vet - name: Test with coverage run: go test -shuffle=on -count=1 -covermode=atomic -coverprofile=coverage.out ./... + - name: Test hardened Hysteria module + working-directory: third_party/hysteria-core + run: go test -shuffle=on -count=1 ./... + - name: Test hardened QUIC core, HTTP/3, and RFC 9002 recovery + working-directory: third_party/quic-go + run: | + go test -shuffle=on -count=1 . -run '^TestServerCancelsConnContextWhenConnectionIDGenerationFails$' + go test -shuffle=on -count=1 ./http3 ./internal/ackhandler + - name: Test Chrome fingerprint, datagrams, and path MTU discovery + working-directory: third_party/quic-go + run: go test -shuffle=on -count=1 ./integrationtests/self -run '^(TestChromeParrot|TestDatagram(Negotiation|SizeLimit)|TestPathMTUDiscovery)' - name: Upload coverage - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4 + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: name: coverage-go-${{ matrix.go }} path: coverage.out @@ -46,24 +61,32 @@ jobs: name: Race detector runs-on: ubuntu-24.04 steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 with: go-version: 1.27.x cache: true - run: make race + - name: Race-test hardened Hysteria resource limits + working-directory: third_party/hysteria-core + run: go test -race -shuffle=on -count=1 ./client ./server ./internal/frag ./internal/protocol ./internal/congestion/... + - name: Race-test hardened QUIC admission and loss recovery + working-directory: third_party/quic-go + run: | + go test -race -shuffle=on -count=1 . -run '^TestServerCancelsConnContextWhenConnectionIDGenerationFails$' + go test -race -shuffle=on -count=1 ./http3 ./internal/ackhandler vulnerability-scan: name: Reachable vulnerability scan runs-on: ubuntu-24.04 steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 with: go-version: 1.27.x cache: true - name: Run govulncheck - run: go run golang.org/x/vuln/cmd/govulncheck@v1.7.0 ./... + run: ./scripts/govulncheck.sh build: name: Build ${{ matrix.goos }}/${{ matrix.goarch }} @@ -74,19 +97,27 @@ jobs: include: - goos: linux goarch: amd64 - output: autocar-linux-amd64 + artifact: autocar-linux-amd64 + binary: autocar + archive: autocar-linux-amd64.tar.gz - goos: linux goarch: arm64 - output: autocar-linux-arm64 + artifact: autocar-linux-arm64 + binary: autocar + archive: autocar-linux-arm64.tar.gz - goos: darwin goarch: arm64 - output: autocar-darwin-arm64 + artifact: autocar-darwin-arm64 + binary: autocar + archive: autocar-darwin-arm64.tar.gz - goos: windows goarch: amd64 - output: autocar-windows-amd64.exe + artifact: autocar-windows-amd64 + binary: autocar.exe + archive: autocar-windows-amd64.zip steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 with: go-version: 1.27.x cache: true @@ -96,19 +127,26 @@ jobs: GOOS: ${{ matrix.goos }} GOARCH: ${{ matrix.goarch }} run: | - mkdir -p dist - go build -trimpath -o "dist/${{ matrix.output }}" ./cmd/autocar - - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4 + mkdir -p dist/package + go build -trimpath -o "dist/package/${{ matrix.binary }}" ./cmd/autocar + cp LICENSE THIRD_PARTY_NOTICES.md dist/package/ + if [ "${{ matrix.goos }}" = "windows" ]; then + (cd dist/package && zip -q -X "../${{ matrix.archive }}" "${{ matrix.binary }}" LICENSE THIRD_PARTY_NOTICES.md) + else + tar -C dist/package -czf "dist/${{ matrix.archive }}" "${{ matrix.binary }}" LICENSE THIRD_PARTY_NOTICES.md + fi + test -s "dist/${{ matrix.archive }}" + - uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: - name: ${{ matrix.output }} - path: dist/${{ matrix.output }} + name: ${{ matrix.artifact }} + path: dist/${{ matrix.archive }} if-no-files-found: error container: name: Multi-platform container build runs-on: ubuntu-24.04 steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Build OCI image index run: | docker buildx create --name autocar-ci --driver docker-container --use diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index ffb416c..0d3ca2c 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -22,17 +22,17 @@ jobs: name: Analyze Go runs-on: ubuntu-24.04 steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 with: go-version: 1.27.x cache: true - - uses: github/codeql-action/init@42947a340483f03ba47bb1a039b2c519aab3df85 # v3 + - uses: github/codeql-action/init@db488ddef3bf6cb639b32c2e9a7c0a7ea8271d28 # v4.37.8 with: languages: go build-mode: manual - name: Build run: go build -o /tmp/autocar-codeql ./cmd/autocar - - uses: github/codeql-action/analyze@42947a340483f03ba47bb1a039b2c519aab3df85 # v3 + - uses: github/codeql-action/analyze@db488ddef3bf6cb639b32c2e9a7c0a7ea8271d28 # v4.37.8 with: category: /language:go diff --git a/.github/workflows/netem.yml b/.github/workflows/netem.yml index 11b1487..af7541e 100644 --- a/.github/workflows/netem.yml +++ b/.github/workflows/netem.yml @@ -19,8 +19,8 @@ jobs: runs-on: ubuntu-24.04 timeout-minutes: 15 steps: - - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 with: go-version: 1.27.x cache: true @@ -32,7 +32,7 @@ jobs: run: sudo env AUTOCAR_ARTIFACT_DIR="$PWD/artifacts/netem" ./scripts/netem-integration.sh "$PWD/bin/autocar" - name: Publish measurements and diagnostics if: always() - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4 + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: name: netem-results path: artifacts/netem diff --git a/Dockerfile b/Dockerfile index dcce9c2..56115ab 100644 --- a/Dockerfile +++ b/Dockerfile @@ -12,6 +12,9 @@ WORKDIR /src RUN apk add --no-cache ca-certificates COPY go.mod go.sum ./ +# Local replace directives are resolved during go mod download, so make the +# audited module forks available before dependency resolution. +COPY third_party ./third_party RUN go mod download COPY . . @@ -23,11 +26,13 @@ RUN CGO_ENABLED=0 GOOS="${TARGETOS}" GOARCH="${TARGETARCH}" \ FROM --platform=$TARGETPLATFORM scratch LABEL org.opencontainers.image.source="https://github.com/cppla/autocar" \ - org.opencontainers.image.description="Authenticated dual-ended QUIC/TLS TCP proxy" \ - org.opencontainers.image.licenses="Apache-2.0" + org.opencontainers.image.description="Secure dual-ended QUIC/TLS TCP and UDP accelerator" \ + org.opencontainers.image.licenses="MIT" COPY --from=build /etc/ssl/certs/ca-certificates.crt /etc/ssl/certs/ca-certificates.crt COPY --from=build --chown=65532:65532 /out/autocar /autocar +COPY --from=build /src/LICENSE /licenses/autocar-LICENSE +COPY --from=build /src/THIRD_PARTY_NOTICES.md /licenses/THIRD_PARTY_NOTICES.md USER 65532:65532 EXPOSE 8443/tcp 8443/udp 1080/tcp 8080/tcp diff --git a/Makefile b/Makefile index a9d5c43..71bd848 100644 --- a/Makefile +++ b/Makefile @@ -8,11 +8,11 @@ LDFLAGS := -s -w \ -X github.com/cppla/autocar/internal/version.Commit=$(COMMIT) \ -X github.com/cppla/autocar/internal/version.Date=$(BUILD_DATE) -.PHONY: all check fmt fmt-check mod-check vet test race build cross-build docker integration-netem clean +.PHONY: all check fmt fmt-check fork-provenance-check mod-check notices notices-check vet test race build cross-build release docker integration-netem clean all: check build -check: fmt-check mod-check vet test +check: fmt-check fork-provenance-check mod-check notices-check vet test fmt: $(GO) fmt ./... @@ -20,12 +20,25 @@ fmt: fmt-check: @test -z "$$(gofmt -l .)" || { gofmt -l .; echo "Go files need formatting" >&2; exit 1; } +fork-provenance-check: + ./scripts/check-fork-provenance.sh + mod-check: $(GO) mod tidy - git diff --exit-code -- go.mod go.sum + cd third_party/hysteria-core && $(GO) mod tidy + cd third_party/quic-go && $(GO) mod tidy + git diff --exit-code -- go.mod go.sum third_party/hysteria-core/go.mod third_party/hysteria-core/go.sum third_party/quic-go/go.mod third_party/quic-go/go.sum + +notices: + $(GO) run ./tools/notices + +notices-check: + $(GO) run ./tools/notices -check vet: $(GO) vet ./... + cd third_party/hysteria-core && $(GO) vet ./... + cd third_party/quic-go && $(GO) vet . ./http3 ./internal/ackhandler test: $(GO) test -shuffle=on -count=1 ./... @@ -37,12 +50,29 @@ build: mkdir -p $(dir $(BINARY)) CGO_ENABLED=0 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o $(BINARY) ./cmd/autocar -cross-build: +cross-build: notices-check mkdir -p dist - CGO_ENABLED=0 GOOS=linux GOARCH=amd64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/autocar-linux-amd64 ./cmd/autocar - CGO_ENABLED=0 GOOS=linux GOARCH=arm64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/autocar-linux-arm64 ./cmd/autocar - CGO_ENABLED=0 GOOS=darwin GOARCH=arm64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/autocar-darwin-arm64 ./cmd/autocar - CGO_ENABLED=0 GOOS=windows GOARCH=amd64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/autocar-windows-amd64.exe ./cmd/autocar + rm -f dist/autocar-linux-amd64 dist/autocar-linux-arm64 dist/autocar-darwin-arm64 dist/autocar-windows-amd64.exe + rm -rf dist/.release-stage-autocar + mkdir -p dist/.release-stage-autocar/autocar-linux-amd64 + mkdir -p dist/.release-stage-autocar/autocar-linux-arm64 + mkdir -p dist/.release-stage-autocar/autocar-darwin-arm64 + mkdir -p dist/.release-stage-autocar/autocar-windows-amd64 + CGO_ENABLED=0 GOOS=linux GOARCH=amd64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/.release-stage-autocar/autocar-linux-amd64/autocar ./cmd/autocar + CGO_ENABLED=0 GOOS=linux GOARCH=arm64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/.release-stage-autocar/autocar-linux-arm64/autocar ./cmd/autocar + CGO_ENABLED=0 GOOS=darwin GOARCH=arm64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/.release-stage-autocar/autocar-darwin-arm64/autocar ./cmd/autocar + CGO_ENABLED=0 GOOS=windows GOARCH=amd64 $(GO) build -trimpath -ldflags "$(LDFLAGS)" -o dist/.release-stage-autocar/autocar-windows-amd64/autocar.exe ./cmd/autocar + cp LICENSE THIRD_PARTY_NOTICES.md dist/.release-stage-autocar/autocar-linux-amd64/ + cp LICENSE THIRD_PARTY_NOTICES.md dist/.release-stage-autocar/autocar-linux-arm64/ + cp LICENSE THIRD_PARTY_NOTICES.md dist/.release-stage-autocar/autocar-darwin-arm64/ + cp LICENSE THIRD_PARTY_NOTICES.md dist/.release-stage-autocar/autocar-windows-amd64/ + tar -C dist/.release-stage-autocar/autocar-linux-amd64 -czf dist/autocar-linux-amd64.tar.gz autocar LICENSE THIRD_PARTY_NOTICES.md + tar -C dist/.release-stage-autocar/autocar-linux-arm64 -czf dist/autocar-linux-arm64.tar.gz autocar LICENSE THIRD_PARTY_NOTICES.md + tar -C dist/.release-stage-autocar/autocar-darwin-arm64 -czf dist/autocar-darwin-arm64.tar.gz autocar LICENSE THIRD_PARTY_NOTICES.md + cd dist/.release-stage-autocar/autocar-windows-amd64 && zip -q -X ../../autocar-windows-amd64.zip autocar.exe LICENSE THIRD_PARTY_NOTICES.md + rm -rf dist/.release-stage-autocar + +release: cross-build docker: docker build --build-arg VERSION="$(VERSION)" --build-arg COMMIT="$(COMMIT)" --build-arg BUILD_DATE="$(BUILD_DATE)" -t autocar:local . diff --git a/README.md b/README.md index 5793a36..570dd34 100644 --- a/README.md +++ b/README.md @@ -1,34 +1,47 @@ -# AutoCAR +

+ AutoCAR — secure dual-ended accelerator +

-AutoCAR 是一个用 Go 编写的双端 TCP 加速代理:本地端提供 SOCKS5、HTTP 和 HTTPS Proxy,远端负责解析域名并连接目标;两端之间优先复用一条经过认证的 QUIC 连接,UDP 不可用时自动切换到 TCP + TLS 1.3。 +

AutoCAR

-> 项目目标是改善高时延、有丢包或大量短连接场景中的体验,而不是承诺“任何网络都更快”。稳定、低时延的直连链路可能更快。请先用内置基准在自己的真实路径上测量。 +

安全、可测量的 Go 双端自适应网络加速器

-## 特性 +

+ CI + CodeQL + netem integration + MIT License +

-- 双端架构:每个代理 TCP 流映射到独立 QUIC 双向流,避免一个应用流的有序传输阻塞其他应用流。 -- 长连接复用:QUIC 会话及其拥塞状态可被后续短连接复用,减少重复传输握手带来的开销。 -- 自动回退:新建流的 QUIC 路径失败后使用独立的 TLS 1.3/TCP 连接,并通过冷却机制避免每条新连接都等待 UDP 超时。 -- 三种本地入口:SOCKS5 CONNECT、HTTP absolute-form HTTP/CONNECT、HTTPS Proxy(代理监听器本身使用 TLS)。 -- 严格安全默认值:TLS 1.3、正常 X.509 主机名校验、无 `skip verify` 开关、禁用 QUIC 0-RTT、可选 mTLS、恒定时间令牌验证。 -- 出口保护:默认阻止环回、私网、链路本地、多播、未指定地址,以及常见 SMTP 提交端口;域名在远端解析并逐个校验,只拨已批准的数字 IP,同时交错竞速 IPv6/IPv4,兼顾 DNS rebinding/SSRF 防护与双栈可用性。 -- 有界资源:协议字段长度、连接数、流数和超时均有限制;中继使用背压而不是无限缓存。 -- 可复现实验:内置上传/下载基准,并提供 Linux `netem` 场景用于直连与隧道的同条件比较。 +AutoCAR 在本地提供 SOCKS5、HTTP 和 HTTPS Proxy,在远端安全地解析域名并连接目标。默认链路采用 Hysteria v2.12.1 的 HTTP/3-over-QUIC 核心;每个方向独立使用真实的 BBRv1,或在双方明确配置带宽后使用 Brutal。UDP 不可达时,新建 TCP 流会自动切换到独立的 TLS 1.3/TCP 回退链路。 -当前版本只代理 TCP。SOCKS5 `BIND`、`UDP ASSOCIATE` 和 QUIC DATAGRAM 尚未实现,收到这些命令会返回标准“不支持”响应。 +它是 split proxy,不把原始 TCP 包套进 UDP,因此不会产生 TCP-over-TCP 的双重可靠传输。目标是改善高 RTT、随机丢包、短连接和多并发场景;稳定、低时延的直连仍可能更快,请始终在真实路径上测量。 -## 数据路径 +## 已实现的加速机制 -```mermaid -flowchart LR - A["应用"] --> P["SOCKS5 / HTTP(S)"] - P --> C["AutoCAR 客户端"] - C -->|"QUIC + TLS 1.3"| S["AutoCAR 服务端"] - C -. "TCP + TLS 1.3 回退" .-> S - S --> D["目标站点"] -``` +| 来源/目标 | AutoCAR 中的实现 | 边界 | +| --- | --- | --- | +| Hysteria v2 | 基于官方 core v2.12.1 的可审计安全加固 fork、HTTP/3 多流、QUIC DATAGRAM、Fast Open、Chrome QUIC 指纹、HTTP/3 cover、可选 Salamander | 不包含实验性的 Gecko、Mimic、端口跳跃或 TUN/TProxy | +| BBR | delivery-rate 与 min-RTT/BDP 模型、pacing,以及 `STARTUP → DRAIN → PROBE_BW → PROBE_RTT`;支持 conservative/standard/aggressive profile | 这是 Hysteria 的 **BBRv1**,不是 Linux 内核 BBRv2/BBRv3 | +| ServerSpeeder/LotServer 的公开目标 | 双端独立发送控制、对端 ACK/RTT/loss 反馈、RFC 9002 packet/time threshold、PTO、热连接拥塞状态复用 | 没有复制 Zeta-TCP 的专有逐包概率算法,也不是内核透明 TCP、FEC 或包复制 | +| 高丢包固定带宽 | 双方协商 `min(发送端上限, 接收端上限)` 后启用 Brutal;根据 ACK/loss 采样补偿并 pacing | 必须显式填准确带宽;会争抢共享链路,默认关闭 | + +默认值是 `bbr + standard`,客户端上下行带宽均为 `0`,且服务端默认忽略客户端带宽提示,所以不会无意启用 Brutal。详细机制、参数和诚实的声明边界见 [加速设计](docs/ACCELERATION.md)。 + +Fast Open 默认关闭;只有显式设置 `--fast-open` 才会让首批应用数据与远端拨号响应重叠。这样能减少一次等待,但目标拒绝等错误可能延迟到第一次读取时才返回。 + +## 代理与安全 + +- SOCKS5:CONNECT 和 UDP ASSOCIATE;UDP 通过 QUIC DATAGRAM 双向传输。 +- HTTP Proxy:absolute-form HTTP 和 CONNECT。 +- HTTPS Proxy:本地代理监听器自身使用 TLS 1.3。 +- 隧道安全:TLS 1.3、正常 X.509 SAN/链验证、强制共享令牌、可选 mTLS;不存在 `skip verify` 开关。 +- 远端出口:域名由服务端解析,每个 TCP/UDP 目标都经过端口、CIDR、特殊用途地址及 DNS rebinding/SSRF 检查,只使用已批准的数字 IP。 +- 资源防护:握手期/已接受的 QUIC 连接、TCP handler、UDP session、双向/单向 stream、HTTP 头和出口 socket 均有硬上限;连接、TCP handler 与 UDP session 还具有跨 QUIC 会话的来源配额(IPv4 地址或 IPv6 `/64`)。TLS/TCP 回退从 accept 到中继结束也有独立的全局与来源连接配额。未认证连接与 TCP 请求头有 deadline,恶意 UDP 分片在分配重组状态前即受限。 +- 可达性:QUIC 失败后对新流使用真实 TLS/TCP,并通过熔断冷却避免 UDP 黑洞造成重复等待。 +- 抗主动探测:默认未认证请求表现为普通 HTTP/3 页面,客户端启用 Chrome QUIC 指纹;受限网络可选择 Salamander 包混淆。 -AutoCAR 是 split proxy,而不是把原始 TCP 包再次塞进 UDP。它在本地终止代理连接、通过 QUIC 流传送字节、再从远端建立新的 TCP 连接,因此不会形成 TCP-over-TCP 或双层可靠重传。 +TLS 保护机密性、完整性和服务端身份;应用仍应使用 HTTPS、SSH 等端到端协议,因为中继知道目标地址,也能看到目标侧明文。网络观察者仍可能看到端点 IP、流量大小和时序。HTTP/3 cover、Salamander 与 TCP 回退提高抗误识别和可达性,但项目不承诺“不可检测”或“永不封锁”。 ## 快速开始 @@ -40,7 +53,7 @@ cd autocar go build -trimpath -o autocar ./cmd/autocar ``` -在服务端生成共享令牌和包含真实域名/IP SAN 的证书: +在服务端生成令牌和包含真实域名/IP SAN 的证书: ```bash ./autocar token --out token @@ -50,19 +63,20 @@ go build -trimpath -o autocar ./cmd/autocar --key server.key ``` -将 `token` 和用于信任的 `server.crt` 通过安全的带外通道复制到客户端。私钥 `server.key` 只留在服务端。 +通过可信带外通道把 `token` 与 `server.crt` 复制到客户端;`server.key` 只留在服务端。启动服务端(UDP 与 TCP 可使用相同端口号): -启动远端。UDP 和 TCP 可以使用同一个端口号: +`autocar cert` 生成与默认 Chrome QUIC 指纹兼容的 ECDSA P-256 证书。若使用外部证书,应选择 ECDSA P-256/P-384 或 RSA;Ed25519 服务端证书需要所有 Hysteria 客户端显式设置 `--disable-chrome-parrot`,否则 TLS 握手会失败并给出提示。 ```bash ./autocar server \ --listen :443 \ + --tcp-listen :443 \ --cert server.crt \ --key server.key \ --token-file token ``` -启动本地端: +启动客户端: ```bash ./autocar client \ @@ -71,74 +85,90 @@ go build -trimpath -o autocar ./cmd/autocar --token-file token ``` -默认监听: +默认入口: | 入口 | 地址 | 示例 | -|---|---:|---| -| SOCKS5 | `127.0.0.1:1080` | `curl --proxy socks5h://127.0.0.1:1080 https://example.com` | +| --- | --- | --- | +| SOCKS5 TCP/UDP | `127.0.0.1:1080` | `curl --proxy socks5h://127.0.0.1:1080 https://example.com` | | HTTP Proxy | `127.0.0.1:8080` | `curl --proxy http://127.0.0.1:8080 https://example.com` | -| HTTPS Proxy | 默认关闭 | 使用 `--https`、`--proxy-cert` 和 `--proxy-key` 开启 | +| HTTPS Proxy | 默认关闭 | 使用 `--https`、`--proxy-cert`、`--proxy-key` 开启 | -`socks5h` 会把域名交给远端解析。HTTP 访问 HTTPS 目标时使用 CONNECT;AutoCAR 不伪造目标证书,也不解密应用到目标站点之间的 HTTPS。 +`socks5h` 会把域名交给远端解析。AutoCAR 不伪造目标证书,也不解密应用到目标站点之间的 HTTPS。 -## 本地代理认证 +## 选择 BBR 或 Brutal -本机独占使用时保留环回默认监听即可。多人机器或非环回监听必须设置本地代理认证: +一般部署直接使用默认 BBR。可按链路偏好选择 profile: ```bash -export AUTOCAR_PROXY_USER=alice -export AUTOCAR_PROXY_PASSWORD='replace-with-a-long-random-secret' +# 共享链路更保守 +./autocar client [其他参数] --bbr-profile conservative -./autocar client \ - --server relay.example.com:443 \ - --ca server.crt \ - --token-file token \ - --socks 127.0.0.1:1080 \ - --http 127.0.0.1:8080 +# Startup 更激进;必须先在自己的链路做公平性与排队延迟测试 +./autocar client [其他参数] --bbr-profile aggressive ``` -SOCKS5 用户名密码和 HTTP Basic 在本地这一跳本身不加密,因此程序默认拒绝把这两个明文入口绑定到非环回地址。跨主机使用请开启带 TLS 的 HTTPS Proxy;只有在已经存在可信外层网络时,才应显式使用 `--allow-public-plaintext`。 +只有已知真实链路容量时才配置 Brutal。官方客户端会在每个方向取声明值与服务端协商上限中的较小值: -## 证书与 mTLS +```bash +# 服务端:每个认证会话最高上传 100 Mbit/s、下载 300 Mbit/s +./autocar server [其他参数] \ + --allow-client-bandwidth \ + --max-upload-mbps 100 \ + --max-download-mbps 300 + +# 客户端:本地链路实测上限 +./autocar client [其他参数] \ + --upload-mbps 80 \ + --download-mbps 250 +``` -客户端必须二选一: +服务端只有显式设置 `--allow-client-bandwidth` 且同时提供两个有限协商上限时才接受 Brutal 提示;默认会强制 BBR/Reno。配置高于实际容量会造成排队、丢包和浪费。上述值是协议协商与 pacing 目标,不是针对恶意客户端的流量整形器;需要不可绕过的限速时,应在主机或云网络层配置 policer。Brutal 不是 Reno/CUBIC 公平模式,共享网络应保留默认 BBR。 -- `--ca `:固定私有 CA/自签名证书,推荐自建部署使用; -- `--system-roots`:明确使用操作系统信任库,适合公共 CA 证书。 +## HTTP/3 cover 与 Salamander -证书名称与连接地址不一致时,用 `--server-name` 指定证书 SAN。程序不会提供跳过验证的选项。若需要双向证书认证,服务端设置 `--client-ca`,客户端同时设置 `--client-cert` 与 `--client-key`。共享令牌仍作为每条隧道流的第二层授权。 +默认模式是标准 HTTP/3 cover:错误令牌或普通探测会得到中性网页,客户端模拟 Chrome QUIC 的可见参数。若 UDP 被按 QUIC 特征干扰,可在两端配置同一条独立强密码: -## 出口策略 +```bash +./autocar token --out obfs-password -默认禁止通过中继访问内网地址,防止被滥用为开放代理或 SSRF 跳板。确实需要访问服务端所在私网时,可在充分信任所有客户端后设置 `--allow-private`;环回、链路本地、多播和未指定地址仍然禁止。默认拒绝端口 `25,465,587`,可用 `--deny-ports` 调整;`--deny-cidrs` 可额外封锁云厂商控制面或部署专用网段。 +./autocar server [其他参数] --obfs-password-file obfs-password +./autocar client [其他参数] --obfs-password-file obfs-password +``` -任何非环回代理监听都必须配置认证。生产环境还应使用主机防火墙,仅向预期客户端开放 UDP/TCP 端口,并优先启用 mTLS。 +也可通过 `AUTOCAR_OBFS_PASSWORD` 提供。Salamander 只是包级混淆,真正的认证与加密仍由 TLS 1.3 完成。启用后线上形态不再是标准 HTTP/3,因此应在“HTTP/3 cover”和“Salamander”之间按网络环境选择,而不是同时宣传两种外观。 -## 性能验证 +## 本地代理认证与 mTLS -基准服务默认仅监听 `127.0.0.1:9000`。只在同一主机测试时可直接启动: +无认证入口只能绑定环回。SOCKS5 用户名密码和 HTTP Basic 在本地这一跳是明文;跨主机使用应开启 HTTPS Proxy,不要把明文入口直接暴露到公网。 ```bash -./autocar bench-server +export AUTOCAR_PROXY_USER=alice +export AUTOCAR_PROXY_PASSWORD='replace-with-a-long-random-secret' + +./autocar client [中继参数] \ + --socks 127.0.0.1:1080 \ + --http 127.0.0.1:8080 ``` -跨主机测量必须显式确认非环回监听: +客户端信任方式必须二选一:`--ca ` 固定私有 CA/自签名证书,或显式使用 `--system-roots`。证书名称与连接地址不一致时设置 `--server-name`。mTLS 使用服务端 `--client-ca` 与客户端 `--client-cert/--client-key`;共享令牌仍保留为第二层授权。 -```bash -./autocar bench-server \ - --listen 0.0.0.0:9000 \ - --allow-public-benchmark -``` +## 出口策略 + +默认拒绝环回、私网、链路本地、多播、未指定地址、IANA 特殊用途地址,以及端口 `25,465,587`。`--allow-private` 仅允许 RFC1918/ULA/CGNAT,仍不会开放环回或云元数据等特殊地址。`--deny-cidrs` 与 `--deny-ports` 可进一步收紧策略。 + +生产环境还应使用主机/云防火墙限制中继 UDP/TCP 端口,给 UDP 设置每源速率与突发上限,并优先启用 mTLS。完整 systemd、容器、防火墙和升级说明见 [部署指南](docs/DEPLOYMENT.md)。 -> `bench-server` 没有认证,远程请求者可以让它持续发送或接收大量数据。 -> 非环回监听只应在临时、受控的测试窗口使用;同时用主机/云防火墙把 -> 端口 9000 严格限制到预期客户端和中继 IP,并在测量后立即停止服务。 -> CLI 默认还把单次传输和并发传输分别限制为 64 MiB 和 16;公开测试时 -> 只应按实际需要调低或谨慎调高 `--max-bytes` / `--max-connections`。 +## 兼容性 -分别测直连和隧道,确保目标、字节数、次数和链路条件完全相同: +当前默认 QUIC wire protocol 是 Hysteria v2.12.1。`--transport=hy2` 与 `--transport=quic` 等价;旧 AutoCAR 自定义 QUIC v1 可临时使用客户端 `--transport=legacy-quic` 配合服务端 `--quic-engine=legacy`。TLS/TCP fallback 继续使用 AutoCAR protocol v1。一次 UDP 端口不能同时运行两种 QUIC wire protocol,升级时必须协调两端或使用不同端口。 + +## 性能验证 + +内置基准会在相同目标、负载和链路条件下比较 direct、hy2/QUIC 与 TLS: ```bash +./autocar bench-server + ./autocar bench-client \ --transport direct \ --target target.example:9000 \ @@ -152,17 +182,18 @@ SOCKS5 用户名密码和 HTTP Basic 在本地这一跳本身不加密,因此 --bytes 8388608 --iterations 7 --warmup 2 --json ``` -最可能受益的是高 RTT、存在随机丢包、多个并发/连续短连接,以及直连路径质量明显差于中继路径的场景。TCP/TLS 回退主要提供可达性,并不声称比直连 TCP 更快。详细方法和 Linux `netem` 脚本见 [基准说明](docs/BENCHMARK.md)。 +Linux `netem` 套件分别验证客户端上传与中继下载在高 RTT/丢包下 BBR 相对 Reno 的收益、Brutal 实际协商/发送、热连接短流收益、错误证书/令牌、UDP→TLS 回退,以及抓包中不存在明文 sentinel: -## 安全边界与封锁 - -TLS 1.3 为客户端到中继的载荷提供机密性、完整性和服务端身份验证;禁用 0-RTT 避免 CONNECT 请求被重放。应用自身使用 HTTPS 时,从应用到目标的内容仍保持端到端加密。 +```bash +make build +sudo ./scripts/netem-integration.sh ./bin/autocar +``` -观察者仍能看到端点 IP、端口、包长、时序,以及使用 UDP/TLS 的事实;中继也知道目标地址。不存在能保证永不被网络运营者识别、限速或封锁的传输。AutoCAR 的策略是提供两个标准、安全的承载路径:优先 QUIC,在 UDP 被阻断时回退到 TLS/TCP,而不是声称“不可检测”。完整威胁模型见 [SECURITY.md](SECURITY.md)。 +`make release` 生成四个平台的发布归档;每个归档都同时包含可执行文件、AutoCAR 的 `LICENSE` 和完整的 `THIRD_PARTY_NOTICES.md`,不会发布缺少许可文件的裸二进制。 -回退只适用于尚在建立或之后新建的代理流。已交付给应用的 QUIC 流若在传输中途失去 UDP 路径,无法安全地把任意 TCP 字节无缝重放到另一条 TLS 连接;该流会失败,由应用重试,随后新流在熔断冷却期内走 TLS。 +CI 中的窄场景速度门只证明被测试的机制有效,不代表所有生产网络都会加速。方法、指标和扩展矩阵见 [基准说明](docs/BENCHMARK.md)。 -## 开发与测试 +## 开发 ```bash go test ./... @@ -170,13 +201,14 @@ go test -race ./... go vet ./... ``` -协议、部署和设计细节分别见: - +- [加速机制与边界](docs/ACCELERATION.md) +- [架构](docs/ARCHITECTURE.md) - [Wire protocol](docs/PROTOCOL.md) -- [Deployment guide](docs/DEPLOYMENT.md) -- [Architecture](docs/ARCHITECTURE.md) -- [Benchmark methodology](docs/BENCHMARK.md) +- [部署指南](docs/DEPLOYMENT.md) +- [基准方法](docs/BENCHMARK.md) +- [安全策略](SECURITY.md) +- [第三方许可](THIRD_PARTY_NOTICES.md) ## License -MIT,见 [LICENSE](LICENSE)。 +AutoCAR 使用 MIT 许可证,见 [LICENSE](LICENSE)。实际链接的全部 Go 依赖及其根级许可、通知和专利声明见自动生成的 [THIRD_PARTY_NOTICES.md](THIRD_PARTY_NOTICES.md)。 diff --git a/SECURITY.md b/SECURITY.md index 92eeb5e..22d4305 100644 --- a/SECURITY.md +++ b/SECURITY.md @@ -28,9 +28,20 @@ and any destination traffic that is itself unencrypted. End-to-end HTTPS remains encrypted between the application and the destination. No transport can guarantee that a network operator will not rate-limit or -block it. AutoCAR provides a standards-compliant TCP/TLS fallback for networks -where UDP is unavailable; it deliberately does not impersonate unrelated -protocols or claim to be undetectable. +block it. AutoCAR's default UDP service is valid HTTP/3 and returns a neutral +cover page to unauthenticated probes; the client uses Hysteria's Chrome QUIC +fingerprint. Optional Salamander changes packet appearance with a separate +pre-shared key, but it is obfuscation rather than encryption. It must never be +treated as a substitute for TLS certificate verification, the relay token or +mTLS. When UDP is unavailable, new TCP flows can use the standards-compliant +TLS/TCP fallback. None of these mechanisms is an undetectability guarantee. + +The default Chrome-parroting ClientHello intentionally follows a signature +scheme list that omits Ed25519. Relay certificates should therefore use ECDSA +P-256/P-384 (`autocar cert` emits P-256) or RSA. An Ed25519 relay certificate is +supported only when the client explicitly uses `--disable-chrome-parrot`; a +matching handshake failure includes this guidance. This switch does not relax +certificate-chain or hostname verification. The relay blocks private, loopback, link-local, multicast, and unspecified destinations by default, and denies common SMTP submission ports. Operators @@ -38,6 +49,71 @@ should keep these defaults unless they fully trust every authenticated client. Deployment-specific control-plane and metadata ranges can be added with `--deny-cidrs`, especially before enabling private destinations. +The same policy is applied to every UDP datagram destination after remote DNS +resolution; only the approved numeric address is used for the actual send. A +logical UDP session remembers at most 256 successfully written numeric +destinations. Once full, it rejects new destinations without evicting existing +ones; failed writes never authorize replies. This bounds memory while retaining +valid delayed-reply filtering semantics. +The local SOCKS5 UDP relay accepts datagrams only from the IP of its associated +TCP control connection. A concrete `UDP ASSOCIATE` address must equal that peer; +a domain is resolved under the dial timeout and must include the peer address. +The relay pins the requested non-zero port, or the first valid source port when +the request uses port zero. SOCKS fragmentation is not reassembled and is +dropped. + +The Hysteria core exposes a deliberately smaller TLS configuration surface than +Go's `tls.Config`. AutoCAR copies server-name/root verification, +`VerifyPeerCertificate` on the client, certificate selection, strict mTLS, and +ECH fields that the core supports. It rejects unsupported security-sensitive +policies before binding or dialing, including `VerifyConnection`, server +`GetConfigForClient`/`VerifyPeerCertificate`, custom verification clocks or +curve policies, custom server ticket handling, and client-authentication modes +other than no certificate or `RequireAndVerifyClientCert`. TLS 1.2-only fields +such as `CipherSuites` and renegotiation are irrelevant to QUIC/TLS 1.3 and do +not cause rejection. The client session cache in the shared CLI TLS config is +used by the per-flow TCP fallback; Hysteria instead keeps a long-lived QUIC +session. + +The relay ignores client-supplied bandwidth hints by default, preventing an +authenticated client from forcing the relay sender into an unbounded Brutal +rate. `--allow-client-bandwidth` is an explicit operator opt-in and is rejected +unless finite upload and download negotiation ceilings are both configured. + +Application limits do not replace host-level denial-of-service controls. The +hardened Hysteria fork caps accepted QUIC sessions globally and per source key, +including cover traffic, after Retry and before handshake allocation. It also +caps active TCP handlers globally and per source key across QUIC connections +before they can wait for a target header. Header +reads have a finite deadline. Unauthenticated HTTP/3 connections must complete +authentication within `--handshake-timeout`, and `--max-uni-streams` gives their +unidirectional control streams a separate small bound. HTTP request headers are +capped at 16 KiB before allocation. UDP session admission is global and shared +per authenticated source key across QUIC connections, and occurs before +allocating defragmentation state; +fragment count and total reassembled bytes are fixed and bounded. +`--max-outbound-tcp` and `--max-outbound-udp` separately cap active target +sockets, while `--max-streams` remains a per-QUIC-connection protocol limit. +Clients behind one NAT share `--max-client-connections`, +`--max-client-fallback-connections`, `--max-client-tcp-handlers`, and +`--max-client-udp-sessions`; the global caps +remain authoritative if a peer can rotate source addresses. +Resource keys are an IPv4 address or a masked IPv6 `/64`, so rotating IPv6 +interface identifiers does not create new buckets. A legitimate NAT or routed +IPv6 `/64` shares its bucket by design. +The relay requires QUIC Retry source-address validation before allocating a +bounded handshake slot. Initial packets still consume kernel/network work, so production +relays should apply firewall rate and burst limits per source, bound file +descriptors and memory with the service manager, and monitor UDP traffic and +authentication failures. + +On the client, connection setup is single-flight and `--max-pending-opens` +bounds stream-open workers whose upstream Hysteria API has no context-aware +variant. Caller deadlines still return immediately; late connections are +closed, and their slot is retained until the underlying call actually exits. + TLS private-key files must be regular files and mode `0600` on Unix. Shared -relay tokens and local-proxy passwords must contain at least 16 bytes; use the -bundled `autocar token` command to generate high-entropy values. +relay tokens, local-proxy passwords and Salamander passwords must contain at +least 16 bytes; use the bundled `autocar token` command to generate independent +high-entropy values. Do not reuse the relay authentication token as the +Salamander password. diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md new file mode 100644 index 0000000..fed26a4 --- /dev/null +++ b/THIRD_PARTY_NOTICES.md @@ -0,0 +1,1019 @@ +# Third-party notices + +> Generated by `go run ./tools/notices`; do not edit by hand. CI verifies this file against the linked build graph. + +AutoCAR itself is licensed under the repository's `LICENSE`. The sections below reproduce every root-level license, notice, and patent file from each non-main Go module reached by `go list -deps -json ./cmd/autocar`. A module-level replacement is recorded so the notice always describes the source that is actually compiled. + +The SHA-256 value is calculated from the upstream file's original bytes; line endings in the displayed copy are normalized for Markdown. + +## `github.com/andybalholm/brotli` `v1.1.0` + +### `LICENSE` + +SHA-256: `3d180008e36922a4e8daec11c34c7af264fed5962d07924aea928c38e8663c94` + +```text +Copyright (c) 2009, 2010, 2013-2016 by the Brotli Authors. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in +all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +THE SOFTWARE. +``` + +## `github.com/apernet/hysteria/core/v2` `v2.12.1` + +Effective source replacement: `./third_party/hysteria-core`. + +### `LICENSE.md` + +SHA-256: `b279cfdac4db4b077f0660b5d8156d50a8bc7bd410036dc356499af43c4e84f5` + +```text +Copyright 2023 Toby + +Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the “Software”), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED “AS IS”, WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. +``` + +## `github.com/apernet/hysteria/extras/v2` `v2.12.1` + +### `LICENSE.md` + +SHA-256: `b279cfdac4db4b077f0660b5d8156d50a8bc7bd410036dc356499af43c4e84f5` + +```text +Copyright 2023 Toby + +Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the “Software”), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED “AS IS”, WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. +``` + +## `github.com/apernet/quic-go` `v0.61.1-0.20260806010916-184d081eef3e` + +Effective source replacement: `./third_party/quic-go`. + +### `LICENSE` + +SHA-256: `77d0b7b53e8abb84cf4dd3f9945a7fdf27044240d2e8023966a721a9a46fe96e` + +```text +MIT License + +Copyright (c) 2016 the quic-go authors & Google, Inc. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +``` + +## `github.com/davecgh/go-spew` `v1.1.1` + +### `LICENSE` + +SHA-256: `1b93a317849ee09d3d7e4f1d20c2b78ddb230b4becb12d7c224c927b9d470251` + +```text +ISC License + +Copyright (c) 2012-2016 Dave Collins + +Permission to use, copy, modify, and/or distribute this software for any +purpose with or without fee is hereby granted, provided that the above +copyright notice and this permission notice appear in all copies. + +THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES +WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF +MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR +ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES +WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN +ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT OF +OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE. +``` + +## `github.com/klauspost/compress` `v1.18.7` + +### `LICENSE` + +SHA-256: `0d9e582ee4bff57bf1189c9e514e6da7ce277f9cd3bc2d488b22fbb39a6d87cf` + +```text +Copyright (c) 2012 The Go Authors. All rights reserved. +Copyright (c) 2019 Klaus Post. All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google Inc. nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +------------------ + +Files: gzhttp/* + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright 2016-2017 The New York Times Company + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. + +------------------ + +Files: s2/cmd/internal/readahead/* + +The MIT License (MIT) + +Copyright (c) 2015 Klaus Post + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. + +--------------------- +Files: snappy/* +Files: internal/snapref/* + +Copyright (c) 2011 The Snappy-Go Authors. All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google Inc. nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +----------------- + +Files: s2/cmd/internal/filepathx/* + +Copyright 2016 The filepathx Authors + +Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. +``` + +## `github.com/pmezard/go-difflib` `v1.0.0` + +### `LICENSE` + +SHA-256: `2eb550be6801c1ea434feba53bf6d12e7c71c90253e0a9de4a4f46cf88b56477` + +```text +Copyright (c) 2013, Patrick Mezard +All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + Redistributions in binary form must reproduce the above copyright +notice, this list of conditions and the following disclaimer in the +documentation and/or other materials provided with the distribution. + The names of its contributors may not be used to endorse or promote +products derived from this software without specific prior written +permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS +IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED +TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A +PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED +TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR +PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF +LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS +SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` + +## `github.com/quic-go/qpack` `v0.6.0` + +### `LICENSE.md` + +SHA-256: `1b6a897efd39b20b3cdce8cd306160d115dbded39d855ceeffe21dc11e4d53df` + +```text +Copyright 2019 Marten Seemann + +Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. +``` + +## `github.com/quic-go/quic-go` `v0.61.0` + +### `LICENSE` + +SHA-256: `77d0b7b53e8abb84cf4dd3f9945a7fdf27044240d2e8023966a721a9a46fe96e` + +```text +MIT License + +Copyright (c) 2016 the quic-go authors & Google, Inc. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +``` + +## `github.com/refraction-networking/utls` `v1.8.2` + +### `LICENSE` + +SHA-256: `2d36597f7117c38b006835ae7f537487207d8ec407aa9d9980794b2030cbc067` + +```text +Copyright (c) 2009 The Go Authors. All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google Inc. nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` + +## `github.com/stretchr/objx` `v0.5.2` + +### `LICENSE` + +SHA-256: `b2663894033a05fd80261176cd8da1d72546e25842d5c1abcc852ca23b6b61b0` + +```text +The MIT License + +Copyright (c) 2014 Stretchr, Inc. +Copyright (c) 2017-2018 objx contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +``` + +## `github.com/stretchr/testify` `v1.11.1` + +### `LICENSE` + +SHA-256: `f8e536c1c7b695810427095dc85f5f80d44ff7c10535e8a9486cf393e2599189` + +```text +MIT License + +Copyright (c) 2012-2020 Mat Ryer, Tyler Bunnell and contributors. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +``` + +## `golang.org/x/crypto` `v0.54.0` + +### `LICENSE` + +SHA-256: `911f8f5782931320f5b8d1160a76365b83aea6447ee6c04fa6d5591467db9dad` + +```text +Copyright 2009 The Go Authors. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google LLC nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` +### `PATENTS` + +SHA-256: `96f408bfae65bf137fc2525d3ecb030271c50c1e90799f87abf8846d8dd505cc` + +```text +Additional IP Rights Grant (Patents) + +"This implementation" means the copyrightable works distributed by +Google as part of the Go project. + +Google hereby grants to You a perpetual, worldwide, non-exclusive, +no-charge, royalty-free, irrevocable (except as stated in this section) +patent license to make, have made, use, offer to sell, sell, import, +transfer and otherwise run, modify and propagate the contents of this +implementation of Go, where such license applies only to those patent +claims, both currently owned or controlled by Google and acquired in +the future, licensable by Google that are necessarily infringed by this +implementation of Go. This grant does not include claims that would be +infringed only as a consequence of further modification of this +implementation. If you or your agent or exclusive licensee institute or +order or agree to the institution of patent litigation against any +entity (including a cross-claim or counterclaim in a lawsuit) alleging +that this implementation of Go or any code incorporated within this +implementation of Go constitutes direct or contributory patent +infringement, or inducement of patent infringement, then any patent +rights granted to you under this License for this implementation of Go +shall terminate as of the date such litigation is filed. +``` + +## `golang.org/x/exp` `v0.0.0-20240506185415-9bf2ced13842` + +### `LICENSE` + +SHA-256: `2d36597f7117c38b006835ae7f537487207d8ec407aa9d9980794b2030cbc067` + +```text +Copyright (c) 2009 The Go Authors. All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google Inc. nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` +### `PATENTS` + +SHA-256: `96f408bfae65bf137fc2525d3ecb030271c50c1e90799f87abf8846d8dd505cc` + +```text +Additional IP Rights Grant (Patents) + +"This implementation" means the copyrightable works distributed by +Google as part of the Go project. + +Google hereby grants to You a perpetual, worldwide, non-exclusive, +no-charge, royalty-free, irrevocable (except as stated in this section) +patent license to make, have made, use, offer to sell, sell, import, +transfer and otherwise run, modify and propagate the contents of this +implementation of Go, where such license applies only to those patent +claims, both currently owned or controlled by Google and acquired in +the future, licensable by Google that are necessarily infringed by this +implementation of Go. This grant does not include claims that would be +infringed only as a consequence of further modification of this +implementation. If you or your agent or exclusive licensee institute or +order or agree to the institution of patent litigation against any +entity (including a cross-claim or counterclaim in a lawsuit) alleging +that this implementation of Go or any code incorporated within this +implementation of Go constitutes direct or contributory patent +infringement, or inducement of patent infringement, then any patent +rights granted to you under this License for this implementation of Go +shall terminate as of the date such litigation is filed. +``` + +## `golang.org/x/net` `v0.57.0` + +### `LICENSE` + +SHA-256: `911f8f5782931320f5b8d1160a76365b83aea6447ee6c04fa6d5591467db9dad` + +```text +Copyright 2009 The Go Authors. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google LLC nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` +### `PATENTS` + +SHA-256: `96f408bfae65bf137fc2525d3ecb030271c50c1e90799f87abf8846d8dd505cc` + +```text +Additional IP Rights Grant (Patents) + +"This implementation" means the copyrightable works distributed by +Google as part of the Go project. + +Google hereby grants to You a perpetual, worldwide, non-exclusive, +no-charge, royalty-free, irrevocable (except as stated in this section) +patent license to make, have made, use, offer to sell, sell, import, +transfer and otherwise run, modify and propagate the contents of this +implementation of Go, where such license applies only to those patent +claims, both currently owned or controlled by Google and acquired in +the future, licensable by Google that are necessarily infringed by this +implementation of Go. This grant does not include claims that would be +infringed only as a consequence of further modification of this +implementation. If you or your agent or exclusive licensee institute or +order or agree to the institution of patent litigation against any +entity (including a cross-claim or counterclaim in a lawsuit) alleging +that this implementation of Go or any code incorporated within this +implementation of Go constitutes direct or contributory patent +infringement, or inducement of patent infringement, then any patent +rights granted to you under this License for this implementation of Go +shall terminate as of the date such litigation is filed. +``` + +## `golang.org/x/sys` `v0.47.0` + +### `LICENSE` + +SHA-256: `911f8f5782931320f5b8d1160a76365b83aea6447ee6c04fa6d5591467db9dad` + +```text +Copyright 2009 The Go Authors. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google LLC nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` +### `PATENTS` + +SHA-256: `96f408bfae65bf137fc2525d3ecb030271c50c1e90799f87abf8846d8dd505cc` + +```text +Additional IP Rights Grant (Patents) + +"This implementation" means the copyrightable works distributed by +Google as part of the Go project. + +Google hereby grants to You a perpetual, worldwide, non-exclusive, +no-charge, royalty-free, irrevocable (except as stated in this section) +patent license to make, have made, use, offer to sell, sell, import, +transfer and otherwise run, modify and propagate the contents of this +implementation of Go, where such license applies only to those patent +claims, both currently owned or controlled by Google and acquired in +the future, licensable by Google that are necessarily infringed by this +implementation of Go. This grant does not include claims that would be +infringed only as a consequence of further modification of this +implementation. If you or your agent or exclusive licensee institute or +order or agree to the institution of patent litigation against any +entity (including a cross-claim or counterclaim in a lawsuit) alleging +that this implementation of Go or any code incorporated within this +implementation of Go constitutes direct or contributory patent +infringement, or inducement of patent infringement, then any patent +rights granted to you under this License for this implementation of Go +shall terminate as of the date such litigation is filed. +``` + +## `golang.org/x/text` `v0.40.0` + +### `LICENSE` + +SHA-256: `911f8f5782931320f5b8d1160a76365b83aea6447ee6c04fa6d5591467db9dad` + +```text +Copyright 2009 The Go Authors. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: + + * Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. + * Redistributions in binary form must reproduce the above +copyright notice, this list of conditions and the following disclaimer +in the documentation and/or other materials provided with the +distribution. + * Neither the name of Google LLC nor the names of its +contributors may be used to endorse or promote products derived from +this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS +"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT +LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR +A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT +OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, +SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT +LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY +THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +``` +### `PATENTS` + +SHA-256: `96f408bfae65bf137fc2525d3ecb030271c50c1e90799f87abf8846d8dd505cc` + +```text +Additional IP Rights Grant (Patents) + +"This implementation" means the copyrightable works distributed by +Google as part of the Go project. + +Google hereby grants to You a perpetual, worldwide, non-exclusive, +no-charge, royalty-free, irrevocable (except as stated in this section) +patent license to make, have made, use, offer to sell, sell, import, +transfer and otherwise run, modify and propagate the contents of this +implementation of Go, where such license applies only to those patent +claims, both currently owned or controlled by Google and acquired in +the future, licensable by Google that are necessarily infringed by this +implementation of Go. This grant does not include claims that would be +infringed only as a consequence of further modification of this +implementation. If you or your agent or exclusive licensee institute or +order or agree to the institution of patent litigation against any +entity (including a cross-claim or counterclaim in a lawsuit) alleging +that this implementation of Go or any code incorporated within this +implementation of Go constitutes direct or contributory patent +infringement, or inducement of patent infringement, then any patent +rights granted to you under this License for this implementation of Go +shall terminate as of the date such litigation is filed. +``` + +## `gopkg.in/yaml.v3` `v3.0.1` + +### `LICENSE` + +SHA-256: `d18f6323b71b0b768bb5e9616e36da390fbd39369a81807cca352de4e4e6aa0b` + +```text + +This project is covered by two different licenses: MIT and Apache. + +#### MIT License #### + +The following files were ported to Go from C files of libyaml, and thus +are still covered by their original MIT license, with the additional +copyright staring in 2011 when the project was ported over: + + apic.go emitterc.go parserc.go readerc.go scannerc.go + writerc.go yamlh.go yamlprivateh.go + +Copyright (c) 2006-2010 Kirill Simonov +Copyright (c) 2006-2011 Kirill Simonov + +Permission is hereby granted, free of charge, to any person obtaining a copy of +this software and associated documentation files (the "Software"), to deal in +the Software without restriction, including without limitation the rights to +use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies +of the Software, and to permit persons to whom the Software is furnished to do +so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. + +### Apache License ### + +All the remaining project files are covered by the Apache license: + +Copyright (c) 2011-2019 Canonical Ltd + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +``` +### `NOTICE` + +SHA-256: `f6c2dd3a67b576eafb89b80200b8b1627230bf3821a0c14cb99a22ac19107d00` + +```text +Copyright 2011-2016 Canonical Ltd. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +``` + +## Brutal code provenance + +AutoCAR's negotiated Brutal controller is provided by the MIT-licensed Hysteria core module identified above. AutoCAR does not vendor, import, or copy the GPL-licensed `tcp-brutal` implementation. Similar terminology describes a traffic-control strategy and does not imply source-code provenance. diff --git a/cmd/autocar/bench.go b/cmd/autocar/bench.go index 659ba00..cc74a6d 100644 --- a/cmd/autocar/bench.go +++ b/cmd/autocar/bench.go @@ -55,14 +55,21 @@ func ensureSafeBenchmarkListener(address string, allowPublic bool) error { } type benchOutput struct { - Mode string `json:"mode"` - Transport string `json:"transport"` - Target string `json:"target"` - Bytes int64 `json:"bytes_per_iteration"` - Iterations int `json:"iterations"` - MedianMbps float64 `json:"median_mbps"` - P95Mbps float64 `json:"p95_mbps"` - Results []float64 `json:"results_mbps"` + Mode string `json:"mode"` + Transport string `json:"transport"` + Acceleration string `json:"acceleration,omitempty"` + NegotiatedTxBytesSec uint64 `json:"negotiated_tx_bytes_per_second,omitempty"` + Target string `json:"target"` + Bytes int64 `json:"bytes_per_iteration"` + Iterations int `json:"iterations"` + MedianMbps float64 `json:"median_mbps"` + P95Mbps float64 `json:"p95_mbps"` + Results []float64 `json:"results_mbps"` +} + +type accelerationReporter interface { + AccelerationMode() string + NegotiatedTx() uint64 } func runBenchClient(parent context.Context, args []string) error { @@ -138,6 +145,10 @@ func runBenchClient(parent context.Context, args []string) error { P95Mbps: p95, Results: results, } + if reporter, ok := dialer.(accelerationReporter); ok { + output.Acceleration = reporter.AccelerationMode() + output.NegotiatedTxBytesSec = reporter.NegotiatedTx() + } if *jsonOutput { encoder := json.NewEncoder(os.Stdout) encoder.SetIndent("", " ") diff --git a/cmd/autocar/common.go b/cmd/autocar/common.go index 76930dd..789dcb4 100644 --- a/cmd/autocar/common.go +++ b/cmd/autocar/common.go @@ -6,6 +6,7 @@ import ( "errors" "flag" "fmt" + "log/slog" "net" "net/netip" "os" @@ -14,41 +15,60 @@ import ( "time" "github.com/cppla/autocar/internal/config" + "github.com/cppla/autocar/internal/hy2" "github.com/cppla/autocar/internal/security" "github.com/cppla/autocar/internal/transport" "github.com/cppla/autocar/internal/tunnel" ) type tunnelFlags struct { - server string - fallback string - mode string - serverName string - caFile string - systemRoots bool - clientCert string - clientKey string - tokenFile string - dialTimeout time.Duration - primaryTimeout time.Duration - openTimeout time.Duration - fallbackTTL time.Duration + server string + fallback string + mode string + serverName string + caFile string + systemRoots bool + clientCert string + clientKey string + tokenFile string + dialTimeout time.Duration + primaryTimeout time.Duration + openTimeout time.Duration + fallbackTTL time.Duration + congestion string + bbrProfile string + uploadMbps uint64 + downloadMbps uint64 + disableLossCompensation bool + fastOpen bool + obfsPasswordFile string + disableChromeParrot bool + maxPendingOpens int } func addTunnelFlags(fs *flag.FlagSet, flags *tunnelFlags) { fs.StringVar(&flags.server, "server", "", "relay host:port (required)") fs.StringVar(&flags.fallback, "fallback-server", "", "TCP/TLS relay host:port; defaults to --server") - fs.StringVar(&flags.mode, "transport", "auto", "transport: auto, quic, or tls") + fs.StringVar(&flags.mode, "transport", "auto", "transport: auto, hy2 (or quic), legacy-quic, or tls") fs.StringVar(&flags.serverName, "server-name", "", "TLS certificate DNS name; defaults to relay host") fs.StringVar(&flags.caFile, "ca", "", "PEM trust anchor for the relay certificate") fs.BoolVar(&flags.systemRoots, "system-roots", false, "trust the operating-system CA set instead of --ca") fs.StringVar(&flags.clientCert, "client-cert", "", "optional mTLS client certificate PEM") fs.StringVar(&flags.clientKey, "client-key", "", "optional mTLS client private key PEM") fs.StringVar(&flags.tokenFile, "token-file", "", "0600 shared-token file; otherwise AUTOCAR_TOKEN") - fs.DurationVar(&flags.dialTimeout, "dial-timeout", 5*time.Second, "QUIC/TLS network dial timeout") + fs.DurationVar(&flags.dialTimeout, "dial-timeout", 5*time.Second, "legacy QUIC/TLS network dial timeout") fs.DurationVar(&flags.primaryTimeout, "quic-attempt-timeout", 5*time.Second, "entire QUIC phase budget before auto-mode TLS fallback") fs.DurationVar(&flags.openTimeout, "open-timeout", 15*time.Second, "overall remote stream open timeout") fs.DurationVar(&flags.fallbackTTL, "fallback-cooldown", 30*time.Second, "time to prefer TLS after a QUIC path failure") + fs.StringVar(&flags.congestion, "congestion", hy2.CongestionBBR, "QUIC congestion controller: bbr or reno; configured bandwidth selects Brutal") + fs.StringVar(&flags.bbrProfile, "bbr-profile", hy2.BBRStandard, "BBR profile: conservative, standard, or aggressive") + fs.Uint64Var(&flags.uploadMbps, "upload-mbps", 0, "known client upload capacity in Mbit/s; nonzero requests negotiated Brutal") + fs.Uint64Var(&flags.downloadMbps, "download-mbps", 0, "known client download capacity in Mbit/s; nonzero requests negotiated Brutal") + fs.BoolVar(&flags.disableLossCompensation, "disable-loss-compensation", false, "disable Brutal ACK/loss-rate compensation") + fs.BoolVar(&flags.fastOpen, "fast-open", false, "return before the exit dial response (lower setup latency, weaker immediate error reporting)") + fs.StringVar(&flags.obfsPasswordFile, "obfs-password-file", "", "0600 Salamander password file; otherwise optional AUTOCAR_OBFS_PASSWORD") + fs.BoolVar(&flags.disableChromeParrot, "disable-chrome-parrot", false, "disable Hysteria's Chrome QUIC fingerprint (diagnostics or Ed25519 relay certificates)") + fs.IntVar(&flags.maxPendingOpens, "max-pending-opens", 256, "maximum in-flight Hysteria TCP stream opens") } type closeDialer interface { @@ -61,12 +81,15 @@ func buildTunnelDialer(flags tunnelFlags) (closeDialer, error) { return nil, errors.New("--server is required") } mode := strings.ToLower(flags.mode) - if mode != "auto" && mode != "quic" && mode != "tls" { - return nil, fmt.Errorf("invalid --transport %q; want auto, quic, or tls", flags.mode) + if mode != "auto" && mode != "hy2" && mode != "quic" && mode != "legacy-quic" && mode != "tls" { + return nil, fmt.Errorf("invalid --transport %q; want auto, hy2, quic, legacy-quic, or tls", flags.mode) } if flags.dialTimeout <= 0 || flags.openTimeout <= 0 { return nil, errors.New("--dial-timeout and --open-timeout must be positive") } + if flags.maxPendingOpens <= 0 || flags.maxPendingOpens > 65536 { + return nil, errors.New("--max-pending-opens must be between 1 and 65536") + } if mode == "auto" && (flags.primaryTimeout <= 0 || flags.primaryTimeout >= flags.openTimeout) { return nil, errors.New("auto mode requires 0 < --quic-attempt-timeout < --open-timeout so TLS fallback retains time") } @@ -122,8 +145,39 @@ func buildTunnelDialer(flags tunnelFlags) (closeDialer, error) { return nil, err } + upload, err := megabitsToBytesPerSecond(flags.uploadMbps) + if err != nil { + return nil, fmt.Errorf("--upload-mbps: %w", err) + } + download, err := megabitsToBytesPerSecond(flags.downloadMbps) + if err != nil { + return nil, fmt.Errorf("--download-mbps: %w", err) + } + obfuscationKey, err := loadOptionalSecret(flags.obfsPasswordFile, "AUTOCAR_OBFS_PASSWORD", 16) + if err != nil { + return nil, fmt.Errorf("load obfuscation password: %w", err) + } + newAcceleratedClient := func() (*hy2.Client, error) { + return hy2.NewClient(hy2.ClientConfig{ + ServerAddress: flags.server, + Token: token, + TLSConfig: tlsConfig, + Congestion: flags.congestion, + BBRProfile: flags.bbrProfile, + MaxTx: upload, + MaxRx: download, + DisableLossCompensation: flags.disableLossCompensation, + FastOpen: flags.fastOpen, + ObfuscationKey: obfuscationKey, + DisableChromeParrot: flags.disableChromeParrot, + MaxPendingOpens: flags.maxPendingOpens, + }) + } + switch mode { - case "quic": + case "hy2", "quic": + return newAcceleratedClient() + case "legacy-quic": return tunnel.NewClient(tunnel.ClientConfig{ ServerAddress: flags.server, Token: token, @@ -144,22 +198,65 @@ func buildTunnelDialer(flags tunnelFlags) (closeDialer, error) { if fallback == "" { fallback = flags.server } - return tunnel.NewClient(tunnel.ClientConfig{ - ServerAddress: flags.server, - FallbackAddress: fallback, - Token: token, - TLSConfig: tlsConfig, - HandshakeTimeout: flags.openTimeout, - QUICDialTimeout: flags.dialTimeout, - PrimaryAttemptTimeout: flags.primaryTimeout, - TLSDialTimeout: flags.dialTimeout, - FallbackCooldown: flags.fallbackTTL, + primary, err := newAcceleratedClient() + if err != nil { + return nil, err + } + fallbackDialer, err := tunnel.NewTLSClient(tunnel.TLSClientConfig{ + ServerAddress: fallback, + Token: token, + TLSConfig: tlsConfig, + HandshakeTimeout: flags.openTimeout, + DialTimeout: flags.dialTimeout, }) + if err != nil { + _ = primary.Close() + return nil, err + } + auto, err := hy2.NewAutoClient(hy2.AutoConfig{ + Primary: primary, + Fallback: fallbackDialer, + AttemptTimeout: flags.primaryTimeout, + Cooldown: flags.fallbackTTL, + OnFallback: func(primaryErr error) { + slog.Warn("QUIC path unavailable; using authenticated TLS fallback", + "error", primaryErr, + "cooldown", flags.fallbackTTL) + }, + }) + if err != nil { + _ = primary.Close() + _ = fallbackDialer.Close() + return nil, err + } + return auto, nil default: panic("unreachable transport mode") } } +func megabitsToBytesPerSecond(value uint64) (uint64, error) { + const bitsPerMegabit = uint64(1_000_000) + if value > ^uint64(0)/bitsPerMegabit { + return 0, errors.New("value is too large") + } + return value * bitsPerMegabit / 8, nil +} + +func loadOptionalSecret(path, envName string, minimumLength int) ([]byte, error) { + if path == "" && os.Getenv(envName) == "" { + return nil, nil + } + value, err := config.LoadSecret(path, envName) + if err != nil { + return nil, err + } + if len(value) < minimumLength { + return nil, fmt.Errorf("secret must be at least %d bytes", minimumLength) + } + return []byte(value), nil +} + func parseDeniedPorts(value string) ([]uint16, error) { value = strings.TrimSpace(value) if value == "" || strings.EqualFold(value, "none") { diff --git a/cmd/autocar/main.go b/cmd/autocar/main.go index 623c901..f263023 100644 --- a/cmd/autocar/main.go +++ b/cmd/autocar/main.go @@ -58,7 +58,7 @@ func run(ctx context.Context, args []string) (err error) { } func printUsage() { - fmt.Fprintln(os.Stderr, `AutoCAR - authenticated dual-ended TCP acceleration + fmt.Fprintln(os.Stderr, `AutoCAR - secure dual-ended TCP/UDP acceleration Usage: autocar server [options] run the remote QUIC/TLS relay diff --git a/cmd/autocar/main_test.go b/cmd/autocar/main_test.go index aeb02a3..eff13ab 100644 --- a/cmd/autocar/main_test.go +++ b/cmd/autocar/main_test.go @@ -2,6 +2,7 @@ package main import ( "context" + "math" "os" "path/filepath" "runtime" @@ -11,6 +12,39 @@ import ( "github.com/cppla/autocar/internal/security" ) +func TestMegabitsToBytesPerSecond(t *testing.T) { + for value, wanted := range map[uint64]uint64{ + 0: 0, + 1: 125_000, + 100: 12_500_000, + } { + got, err := megabitsToBytesPerSecond(value) + if err != nil || got != wanted { + t.Fatalf("%d Mbit/s = %d B/s, %v; want %d", value, got, err, wanted) + } + } + if _, err := megabitsToBytesPerSecond(math.MaxUint64); err == nil { + t.Fatal("overflowing bandwidth was accepted") + } +} + +func TestLoadOptionalSecret(t *testing.T) { + t.Setenv("AUTOCAR_TEST_OPTIONAL_SECRET", "") + value, err := loadOptionalSecret("", "AUTOCAR_TEST_OPTIONAL_SECRET", 16) + if err != nil || value != nil { + t.Fatalf("empty optional secret = %q, %v", value, err) + } + t.Setenv("AUTOCAR_TEST_OPTIONAL_SECRET", "short") + if _, err := loadOptionalSecret("", "AUTOCAR_TEST_OPTIONAL_SECRET", 16); err == nil { + t.Fatal("short optional secret was accepted") + } + t.Setenv("AUTOCAR_TEST_OPTIONAL_SECRET", strings.Repeat("x", 16)) + value, err = loadOptionalSecret("", "AUTOCAR_TEST_OPTIONAL_SECRET", 16) + if err != nil || string(value) != strings.Repeat("x", 16) { + t.Fatalf("optional secret = %q, %v", value, err) + } +} + func TestParseDeniedPorts(t *testing.T) { ports, err := parseDeniedPorts("25, 443,25") if err != nil { @@ -175,6 +209,18 @@ func TestRunHelpAndUnknownCommand(t *testing.T) { } } +func TestServerRejectsFallbackSourceLimitAboveGlobalLimit(t *testing.T) { + err := runServer(context.Background(), []string{ + "--cert", "unused.crt", + "--key", "unused.key", + "--max-streams", "1", + "--max-client-fallback-connections", "2", + }) + if err == nil || !strings.Contains(err.Error(), "--max-client-fallback-connections") { + t.Fatalf("unexpected error: %v", err) + } +} + func TestPercentile(t *testing.T) { values := []float64{1, 2, 3, 4, 5} if got := percentile(values, 0.5); got != 3 { diff --git a/cmd/autocar/server.go b/cmd/autocar/server.go index 2c8eb96..c6c8926 100644 --- a/cmd/autocar/server.go +++ b/cmd/autocar/server.go @@ -14,6 +14,7 @@ import ( "time" "github.com/cppla/autocar/internal/config" + "github.com/cppla/autocar/internal/hy2" "github.com/cppla/autocar/internal/security" "github.com/cppla/autocar/internal/tunnel" ) @@ -21,6 +22,7 @@ import ( func runServer(parent context.Context, args []string) error { fs := flag.NewFlagSet("server", flag.ContinueOnError) listen := fs.String("listen", ":443", "QUIC UDP listen address") + quicEngine := fs.String("quic-engine", "hy2", "UDP engine: hy2 or legacy") tcpListen := fs.String("tcp-listen", "", "TLS/TCP fallback address; defaults to --listen") disableFallback := fs.Bool("disable-tcp-fallback", false, "disable the TCP/TLS fallback listener") certFile := fs.String("cert", "", "server certificate PEM (required)") @@ -30,10 +32,27 @@ func runServer(parent context.Context, args []string) error { allowPrivate := fs.Bool("allow-private", false, "allow RFC1918/ULA/CGNAT destinations (loopback remains blocked)") deniedPortsText := fs.String("deny-ports", "25,465,587", "comma-separated denied destination ports, or none") deniedCIDRsText := fs.String("deny-cidrs", "", "additional comma-separated denied destination CIDRs/IPs") - maxStreams := fs.Int("max-streams", 1024, "maximum concurrent streams per transport listener") - maxConnections := fs.Int("max-connections", 256, "maximum accepted QUIC connections") + maxStreams := fs.Int("max-streams", 1024, "maximum incoming streams per QUIC connection (legacy/TLS use a listener-wide bound)") + maxUniStreams := fs.Int("max-uni-streams", 8, "maximum incoming unidirectional streams per Hysteria QUIC connection") + maxConnections := fs.Int("max-connections", 256, "global maximum accepted QUIC sessions, including unauthenticated cover traffic") + maxClientConnections := fs.Int("max-client-connections", 32, "maximum accepted QUIC sessions per source IPv4 or IPv6 /64 across authenticated and cover traffic") + maxClientFallbackConnections := fs.Int("max-client-fallback-connections", 0, "maximum TLS fallback connections per source IPv4 or IPv6 /64; zero uses min(32, --max-streams)") + maxOutboundTCP := fs.Int("max-outbound-tcp", 1024, "global maximum active Hysteria exit TCP connections") + maxOutboundUDP := fs.Int("max-outbound-udp", 256, "global maximum active Hysteria UDP sessions") + maxClientTCPHandlers := fs.Int("max-client-tcp-handlers", 128, "maximum Hysteria TCP handlers per source IPv4 or IPv6 /64 across QUIC connections") + maxClientUDPSessions := fs.Int("max-client-udp-sessions", 64, "maximum Hysteria UDP sessions per authenticated source IPv4 or IPv6 /64 across QUIC connections") + congestion := fs.String("congestion", hy2.CongestionBBR, "QUIC congestion controller: bbr or reno") + bbrProfile := fs.String("bbr-profile", hy2.BBRStandard, "BBR profile: conservative, standard, or aggressive") + maxUploadMbps := fs.Uint64("max-upload-mbps", 0, "maximum negotiated Brutal upload target in Mbit/s; zero is automatic BBR") + maxDownloadMbps := fs.Uint64("max-download-mbps", 0, "maximum negotiated Brutal download target in Mbit/s; zero is automatic BBR") + allowClientBandwidth := fs.Bool("allow-client-bandwidth", false, "explicitly allow finite client hints to negotiate Brutal (requires both server ceilings)") + disableLossCompensation := fs.Bool("disable-loss-compensation", false, "disable Brutal ACK/loss-rate compensation") + disableUDP := fs.Bool("disable-udp", false, "disable QUIC DATAGRAM and SOCKS5 UDP ASSOCIATE") + udpIdleTimeout := fs.Duration("udp-idle-timeout", 60*time.Second, "idle timeout for each UDP association") + obfsPasswordFile := fs.String("obfs-password-file", "", "0600 Salamander password file; otherwise optional AUTOCAR_OBFS_PASSWORD") + masqueradeName := fs.String("masquerade-name", "", "neutral site name returned to unauthenticated HTTP/3 probes") dialTimeout := fs.Duration("dial-timeout", 4*time.Second, "remote destination dial timeout") - handshakeTimeout := fs.Duration("handshake-timeout", 10*time.Second, "per-stream authentication/open timeout") + handshakeTimeout := fs.Duration("handshake-timeout", 10*time.Second, "authentication and initial stream-open timeout") if err := fs.Parse(args); err != nil { return err } @@ -43,9 +62,37 @@ func runServer(parent context.Context, args []string) error { if *maxStreams <= 0 { return errors.New("--max-streams must be positive") } + if *maxUniStreams < 3 || *maxUniStreams > 1024 { + return errors.New("--max-uni-streams must be between 3 and 1024") + } if *maxConnections <= 0 { return errors.New("--max-connections must be positive") } + if *maxClientConnections <= 0 || *maxClientConnections > *maxConnections { + return errors.New("--max-client-connections must be positive and no greater than --max-connections") + } + if !*disableFallback && (*maxClientFallbackConnections < 0 || *maxClientFallbackConnections > *maxStreams) { + return errors.New("--max-client-fallback-connections must be zero or positive and no greater than --max-streams") + } + if *maxOutboundTCP <= 0 || *maxOutboundUDP <= 0 { + return errors.New("--max-outbound-tcp and --max-outbound-udp must be positive") + } + if *maxClientTCPHandlers <= 0 || *maxClientTCPHandlers > *maxOutboundTCP { + return errors.New("--max-client-tcp-handlers must be positive and no greater than --max-outbound-tcp") + } + if *maxClientUDPSessions <= 0 || *maxClientUDPSessions > *maxOutboundUDP { + return errors.New("--max-client-udp-sessions must be positive and no greater than --max-outbound-udp") + } + if *allowClientBandwidth && (*maxUploadMbps == 0 || *maxDownloadMbps == 0) { + return errors.New("--allow-client-bandwidth requires nonzero --max-upload-mbps and --max-download-mbps") + } + if *udpIdleTimeout < 2*time.Second || *udpIdleTimeout > 10*time.Minute { + return errors.New("--udp-idle-timeout must be between 2s and 10m") + } + engine := strings.ToLower(strings.TrimSpace(*quicEngine)) + if engine != "hy2" && engine != "legacy" { + return errors.New("--quic-engine must be hy2 or legacy") + } if *tcpListen == "" { *tcpListen = *listen } @@ -90,18 +137,74 @@ func runServer(parent context.Context, args []string) error { }, }) - quicServer, err := tunnel.ListenQUIC(tunnel.QUICServerConfig{ - Address: *listen, - Token: token, - TLSConfig: tlsConfig, - Dialer: safeDialer, - HandshakeTimeout: *handshakeTimeout, - DialTimeout: *dialTimeout, - MaxConcurrentStreams: *maxStreams, - MaxConnections: *maxConnections, - }) - if err != nil { - return err + type udpRelay interface { + Addr() net.Addr + Serve(context.Context) error + Close() error + } + var quicServer udpRelay + if engine == "legacy" { + if *obfsPasswordFile != "" || os.Getenv("AUTOCAR_OBFS_PASSWORD") != "" { + return errors.New("Salamander obfuscation requires --quic-engine=hy2") + } + legacyServer, listenErr := tunnel.ListenQUIC(tunnel.QUICServerConfig{ + Address: *listen, + Token: token, + TLSConfig: tlsConfig, + Dialer: safeDialer, + HandshakeTimeout: *handshakeTimeout, + DialTimeout: *dialTimeout, + MaxConcurrentStreams: *maxStreams, + MaxConnections: *maxConnections, + }) + if listenErr != nil { + return listenErr + } + quicServer = legacyServer + } else { + maxUpload, conversionErr := megabitsToBytesPerSecond(*maxUploadMbps) + if conversionErr != nil { + return fmt.Errorf("--max-upload-mbps: %w", conversionErr) + } + maxDownload, conversionErr := megabitsToBytesPerSecond(*maxDownloadMbps) + if conversionErr != nil { + return fmt.Errorf("--max-download-mbps: %w", conversionErr) + } + obfuscationKey, secretErr := loadOptionalSecret(*obfsPasswordFile, "AUTOCAR_OBFS_PASSWORD", 16) + if secretErr != nil { + return fmt.Errorf("load obfuscation password: %w", secretErr) + } + acceleratedServer, listenErr := hy2.Listen(hy2.ServerConfig{ + Address: *listen, + Token: token, + TLSConfig: tlsConfig, + Dialer: safeDialer, + Congestion: *congestion, + BBRProfile: *bbrProfile, + MaxTx: maxDownload, + MaxRx: maxUpload, + AllowClientBandwidth: *allowClientBandwidth, + DisableLossCompensation: *disableLossCompensation, + DisableUDP: *disableUDP, + ObfuscationKey: obfuscationKey, + UDPIdleTimeout: *udpIdleTimeout, + DialTimeout: *dialTimeout, + MaxConcurrentStreams: *maxStreams, + MaxIncomingUniStreams: *maxUniStreams, + MaxConnections: *maxConnections, + MaxClientConnections: *maxClientConnections, + MaxOutboundTCP: *maxOutboundTCP, + MaxOutboundUDP: *maxOutboundUDP, + MaxClientTCPHandlers: *maxClientTCPHandlers, + MaxClientUDPSessions: *maxClientUDPSessions, + TCPRequestTimeout: *handshakeTimeout, + AuthenticationTimeout: *handshakeTimeout, + MasqueradeHandler: hy2.NewCoverHandler(*masqueradeName), + }) + if listenErr != nil { + return listenErr + } + quicServer = acceleratedServer } defer quicServer.Close() @@ -115,6 +218,7 @@ func runServer(parent context.Context, args []string) error { HandshakeTimeout: *handshakeTimeout, DialTimeout: *dialTimeout, MaxConcurrentStreams: *maxStreams, + MaxClientConnections: *maxClientFallbackConnections, }) if err != nil { return err @@ -140,7 +244,7 @@ func runServer(parent context.Context, args []string) error { } }() } - start("quic", quicServer.Addr().String(), quicServer.Serve) + start(engine, quicServer.Addr().String(), quicServer.Serve) if tlsServer != nil { start("tls", tlsServer.Addr().String(), tlsServer.Serve) } diff --git a/docs/ACCELERATION.md b/docs/ACCELERATION.md new file mode 100644 index 0000000..690176a --- /dev/null +++ b/docs/ACCELERATION.md @@ -0,0 +1,195 @@ +# Acceleration design + +AutoCAR combines three public design families without claiming to be a drop-in +replacement for any of them: + +| Design family | What AutoCAR adopts | What AutoCAR does not claim | +| --- | --- | --- | +| Hysteria v2 | HTTP/3 over QUIC, a persistent multiplexed session, Fast Open, negotiated Brutal, QUIC DATAGRAM, Chrome-oriented handshake shaping, HTTP/3 cover handling, and optional Salamander | Port hopping, Mimic, a user-facing ECH setup, or invisibility | +| BBR | A real userspace BBRv1-derived delivery-rate/minimum-RTT model, BDP-based pacing and congestion window, and the four BBR phases | Linux kernel TCP BBR, BBRv2, or BBRv3 | +| ServerSpeeder/LotServer/Zeta-TCP objectives | ACK-driven feedback in both directions, paced sending, warm state, standard early loss detection, PTO probes, and independent multiplexed streams | Proprietary prediction, redundant retransmission, FEC, transparent TCP interception, or protocol compatibility | + +The implementation comes from an in-tree, security-hardened fork of the pinned +MIT-licensed Hysteria v2.12.1 core and its QUIC fork. AutoCAR adds admission +before TCP handlers and UDP defragmentation state, finite fragment bounds, and +exit policy around it. It does not copy the GPL `tcp-brutal` project. See +[the patch record](../third_party/hysteria-core/AUTOCAR_PATCHES.md) and +[THIRD_PARTY_NOTICES.md](../THIRD_PARTY_NOTICES.md). + +## BBR mode (the default) + +Leaving client `--upload-mbps=0 --download-mbps=0` selects the configured +model controller. `--congestion=bbr --bbr-profile=standard` is the default on +both endpoints. + +This is a real BBRv1-derived sender, not the old quic-go default controller +under a different name. For traffic sent by each endpoint it: + +1. samples delivered bytes over send/ACK intervals to estimate bottleneck + delivery rate; +2. tracks the minimum observed RTT, expiring the estimate periodically so the + path can be remeasured; +3. derives the bandwidth-delay product, approximately + `delivery_rate * min_rtt`; +4. applies a pacing gain to the estimated delivery rate; and +5. bounds in-flight data with a congestion-window gain around the BDP, with + loss-recovery limits when packets are declared lost. + +The BBR state machine is: + +| Phase | Purpose | +| --- | --- | +| STARTUP | Increase pacing quickly while delivery bandwidth continues to grow | +| DRAIN | Pace below the estimate to remove the queue accumulated during STARTUP | +| PROBE_BW | Cycle pacing gains around the bandwidth estimate to look for new capacity while controlling the queue | +| PROBE_RTT | Temporarily reduce in-flight data to refresh the minimum-RTT model | + +In the pinned implementation, the PROBE_BW gain cycle is `1.25, 0.75, 1, 1, +1, 1, 1, 1`. The minimum-RTT sample expires after 10 seconds; PROBE_RTT lasts +at least 200 ms after the in-flight target is reached. These are implementation +details of the pinned version and may change only with an explicit dependency +upgrade and review. + +### Profiles + +Profiles change how quickly BBR probes and how conservatively it handles an +overshot path. They do not change the Hysteria wire protocol. + +| Profile | STARTUP pacing gain | STARTUP CWND gain | Steady CWND gain | Growth rounds | Intended use | +| --- | ---: | ---: | ---: | ---: | --- | +| `conservative` | 2.25 | 1.75 | 1.75 | 2 | Shallow buffers, shared access links, or latency-sensitive paths; enables drain-to-target, overshoot detection, and estimate safeguards | +| `standard` | 2.885 | 2.0 | 2.0 | 3 | General default | +| `aggressive` | 3.0 | 2.25 | 2.5 | 4 | Controlled high-BDP paths; allows more startup ACK aggregation and queue pressure | + +Choose the profile independently on the client and relay because each setting +controls only that endpoint's sender. `--congestion=reno` is available as a +diagnostic/fairness baseline. Client bandwidth hints are ignored by default. +They can select Brutal only after the relay operator explicitly enables +`--allow-client-bandwidth` with finite ceilings in both directions. + +BBR is model-based, not magic. A bad route, insufficient relay capacity, +policing, CPU saturation, or an already optimal direct route can erase any +benefit. BBRv1 can also compete aggressively with loss-based flows and can +build queues on paths where its model is inaccurate. + +## Brutal mode (explicit bandwidth only) + +Brutal is selected direction by direction when the client provides a non-zero +capacity and the relay explicitly allows client bandwidth with two finite +ceilings. AutoCAR deliberately has no "guess a large number" default. + +| Traffic direction | Client hint | Relay negotiation ceiling | +| --- | --- | --- | +| Client to relay / upload | `--upload-mbps` | `--max-upload-mbps` | +| Relay to client / download | `--download-mbps` | `--max-download-mbps` | + +The relay opt-in rejects a zero ceiling. A zero client hint means unknown +capacity and therefore keeps BBR/Reno for that direction. With two non-zero +values, the lower value wins. The negotiated value is a sender pacing target, +not a throughput guarantee. + +Example for a measured 20 Mbit/s upload and 100 Mbit/s download: + +```sh +# Relay policy for each authenticated client +autocar server [server options] \ + --allow-client-bandwidth \ + --max-upload-mbps=20 \ + --max-download-mbps=100 + +# Client's measured access-link capacities +autocar client [client options] \ + --upload-mbps=20 \ + --download-mbps=100 +``` + +The sender keeps five one-second ACK/loss sample slots. After at least 50 +packet samples, it computes `ack_rate = ACKed / (ACKed + lost)` and clamps the +rate to a minimum of `0.8`. Pacing is approximately +`negotiated_rate / ack_rate`; therefore loss compensation is capped at about +`1 / 0.8 = 1.25x`. Its congestion window is approximately two smoothed RTTs +of that compensated rate. `--disable-loss-compensation` fixes the ACK rate at +one; set it on both endpoints if compensation must be disabled in both +directions. + +Brutal intentionally keeps sending near the declared rate instead of backing +off like a conventional congestion-fair controller. It can harm other users, +trigger policers, and waste bandwidth when the entered value exceeds the real +bottleneck. Use it only on a link you control or have permission to reserve, +enter a conservative measured capacity, and configure relay negotiation +ceilings. These values are not a non-bypassable traffic policer; use host or +cloud shaping for hard limits. Keep the zero-bandwidth BBR default on shared or +unknown networks. + +## Loss recovery and dual-ended feedback + +Congestion control and retransmission are separate layers. BBR and Brutal +consume the same QUIC ACK/loss events; neither replaces QUIC loss detection. +The pinned QUIC transport follows RFC 9002 with: + +- packet-threshold loss after three newer packet numbers are acknowledged; +- time-threshold loss at 9/8 of the relevant RTT estimate; and +- probe timeout (PTO) packets with exponential backoff when acknowledgements + stop arriving. + +Both client and relay are QUIC senders and receivers. ACKs flowing in each +direction continuously return RTT, delivery, and loss observations to the +opposite sender. That is the concrete dual-ended feedback mechanism behind +AutoCAR's "reverse-control" goal. It is auditable standard QUIC behavior, not +an assertion that AutoCAR reconstructed Zeta-TCP's private algorithm. + +QUIC retransmits lost reliable stream frames, but does not retransmit QUIC +DATAGRAM payloads. AutoCAR adds no speculative retransmission or FEC. Adding +redundancy without a measured policy could amplify congestion and would +require a separate protocol and fairness review. + +## Short-flow and multiplexing gains + +While the current long-lived QUIC session remains connected, it keeps TLS, +RTT, path-MTU, and controller state warm. Each new TCP proxy flow opens a +stream instead of a new end-to-end TCP connection between the AutoCAR +endpoints. A reconnect creates a fresh session and therefore starts cold; no +TLS resumption or congestion/PMTU state is claimed across reconnects. Fast +Open is disabled by default; when explicitly enabled, it lets the first bytes +be written before the exit-dial response reaches the client. + +These mechanisms are most visible for sequential short operations on a +high-RTT path. Independent QUIC streams also prevent a lost ordered byte in one +logical flow from imposing TCP-style application head-of-line blocking on all +other logical flows. They do not remove propagation delay or make the final +relay-to-destination TCP handshake disappear. + +## Hysteria traffic-shaping features + +- **HTTP/3 cover:** without packet obfuscation, unauthenticated requests see a + neutral HTTP/3 service rather than a distinctive tunnel error. +- **Chrome parrot:** enabled by default, it selects Hysteria/quic-go handshake + traits including Chrome-oriented connection-ID behavior. It is a fingerprint + reduction, not proof that all traffic is identical to a browser. Its + Chrome-compatible signature list requires an ECDSA P-256/P-384 or RSA relay + certificate; Ed25519 requires `--disable-chrome-parrot` on the client. +- **Salamander:** an optional shared secret wraps UDP packets before QUIC. It + changes the observable packet form, so normal HTTP/3 cover probing is no + longer available in that mode. +- **TLS/TCP fallback:** `auto` gives new TCP flows a real encrypted TCP path + when UDP is unavailable. UDP associations have no TCP fallback. + +AutoCAR does not currently expose port hopping, Hysteria Mimic, or ECH +provisioning. It cannot promise resistance to endpoint blocking, statistical +traffic analysis, global observation, or traffic-volume correlation. + +## How to verify a deployment + +Run the repository's netem suite first, then repeat the matrix in +[BENCHMARK.md](BENCHMARK.md) on the intended route. At minimum compare: + +1. direct, BBR `conservative`, BBR `standard`, and BBR `aggressive` with both + bandwidth hints zero; +2. Brutal with truthful capacities and relay caps; +3. download and upload, short and bulk payloads, and concurrency greater than + one; and +4. clean, delayed, lossy, and reordered path profiles. + +Retain raw results, packet captures without payload secrets, CPU data, and the +exact build/configuration. A result from one narrow CI profile is evidence for +that profile only, not a universal acceleration claim. diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index fee0f52..29a8c42 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -1,66 +1,159 @@ # Architecture -AutoCAR is a split proxy. It terminates a local proxy connection, carries its -byte stream over an authenticated tunnel, and creates a new TCP connection at -the relay. It never encapsulates raw TCP segments, avoiding nested TCP -retransmission and nested congestion-control loops. +AutoCAR is a split proxy, not a kernel TCP optimizer. The client terminates a +local SOCKS5 or HTTP(S) proxy request, carries it through an authenticated +dual-ended tunnel, and the relay creates a new TCP or UDP flow to the +destination. Raw TCP segments are never nested inside another TCP stream. ```mermaid -flowchart LR - A["Application"] --> P["SOCKS5 / HTTP(S)"] - P --> C["Tunnel client"] - C -->|"QUIC streams"| R["Authenticated relay"] - C -. "TLS/TCP fallback" .-> R +flowchart TB + A["Application"] --> P["SOCKS5 / HTTP(S) proxy"] + P --> C["AutoCAR client"] + C -->|"Hysteria v2 over HTTP/3 + QUIC"| R["AutoCAR relay"] + C -. "TLS 1.3 / TCP fallback" .-> R R --> D["Destination"] ``` -## Design principles - -1. **Standard cryptography.** TLS 1.3, normal X.509 validation, optional mTLS, - and no custom cipher or certificate-verification bypass. -2. **One flow, one stream.** Each proxied TCP connection maps to an independent - QUIC bidirectional stream so loss in one ordered stream does not impose - application-level head-of-line blocking on every other stream. -3. **Bounded parsing.** Every variable-size protocol field has a small hard - limit before allocation. Timeouts, stream limits, and backpressure bound - resource use. -4. **Remote resolution with egress policy.** Hostnames are preserved across - the tunnel. The relay resolves them, filters every resulting address, then - dials only approved numeric IPs to prevent DNS-rebinding bypasses. Approved - IPv6 and IPv4 candidates are interleaved and staggered so one blackholed - address family does not consume the whole destination timeout. -5. **Correctness before aggression.** The default QUIC congestion controller - is used until reproducible measurements justify a maintained transport - fork. The TCP fallback favors reachability, not acceleration. -6. **Explicit limitations.** AutoCAR cannot remove propagation delay or - guarantee higher throughput on every path. Stable, clean direct paths may - be faster because a relay adds work and distance. - -## Data plane - -The QUIC client maintains a warm authenticated connection. Opening a proxy -flow creates a bidirectional QUIC stream and sends a bounded request containing -the command, destination, and authentication proof. The relay validates the -request and destination policy before dialing. A success response switches the -stream into opaque byte-forwarding mode. Half-close is propagated in both -directions. - -When QUIC cannot be established, the optional fallback opens one TLS 1.3 TCP -connection per proxied flow and uses the same request/response framing. This -avoids building an unsafe custom multiplexer over a single TCP byte stream, -but it does not provide QUIC's stream independence. - -## Control plane - -Configuration is local and file/flag based. Tokens should be supplied through -an environment variable or a permission-restricted file rather than a command -line visible to other users. Certificates can be public-CA certificates or a -private certificate generated by the bundled command and distributed out of -band. - -## Future work - -- SOCKS5 UDP ASSOCIATE over QUIC DATAGRAM. -- Separate interactive and bulk QUIC sessions with bounded fair scheduling. -- Reproducible evaluation of CUBIC or BBRv3-based QUIC congestion control. -- Multi-relay path measurement and policy-based selection. +## Components + +| Component | Responsibility | +| --- | --- | +| Local proxy | SOCKS5 CONNECT and UDP ASSOCIATE, HTTP absolute-form requests, HTTP CONNECT, and an optional TLS-protected HTTPS proxy listener | +| Hysteria adapter | Reconnecting Hysteria v2.12.1 client/server, HTTP/3 authentication, TCP streams, QUIC DATAGRAM, Fast Open, BBR/Brutal selection, cover handling, and optional Salamander wrapping | +| Automatic dialer | Tries Hysteria v2 over UDP, opens a bounded circuit breaker after a path failure, and sends new TCP flows through the real TLS/TCP fallback | +| Legacy tunnel | Preserves AutoCAR wire protocol v1 over legacy QUIC and over the TLS/TCP fallback | +| Safe outbound | Resolves names at the relay, rejects unsafe results, dials approved numeric addresses, and rechecks every UDP destination | + +The default UDP engine is `hy2`. On the client, `--transport=quic` is an alias +for `--transport=hy2`; it no longer selects AutoCAR's original QUIC protocol. +The old engine remains available as `--transport=legacy-quic` together with +server `--quic-engine=legacy`. See [PROTOCOL.md](PROTOCOL.md) for the exact +compatibility matrix. + +## Dual-ended acceleration + +One long-lived QUIC connection carries many independent flows. A proxied TCP +connection maps to one bidirectional QUIC stream; a SOCKS5 UDP association maps +to a Hysteria UDP session carried by QUIC DATAGRAM frames. This preserves a +warm RTT and congestion model across short flows while that session remains +connected, and avoids TCP's +connection-wide application head-of-line blocking between unrelated streams. +Reconnection creates a cold QUIC/TLS/controller/PMTU session; state is not +claimed to survive it. + +Each endpoint controls the traffic it sends. QUIC ACKs provide delivery, +loss, and RTT feedback from the opposite endpoint, so both client-to-relay and +relay-to-client directions adapt independently: + +- with bandwidth hints left at zero, each sender uses real BBRv1-derived model + control by default (or Reno when explicitly selected); +- with a non-zero client bandwidth hint, the corresponding direction + negotiates a capped Brutal pacing rate with the relay; and +- QUIC loss recovery remains the standards-based RFC 9002 packet-threshold, + time-threshold, and probe-timeout machinery regardless of congestion mode. + +This ACK feedback satisfies the public objective commonly called +"reverse-control" or dual-ended feedback. It is not the proprietary prediction +or retransmission algorithm from ServerSpeeder/LotServer/Zeta-TCP, and AutoCAR +does not claim protocol compatibility with those products. The controller +details and boundaries are documented in [ACCELERATION.md](ACCELERATION.md). + +## TCP path + +The Hysteria client maintains a reconnecting authenticated HTTP/3 session. +Opening a proxy flow creates a bidirectional stream and sends a bounded target +address request. Fast Open is disabled by default. When explicitly enabled, +the stream can accept the first application bytes before the relay's +destination-dial response is read; a refusal is surfaced on the first read. +Stream limits, open deadlines, QUIC flow control, and operating-system +backpressure bound resource use. `MaxIncomingStreams` applies per QUIC +connection; the hardened core also applies listener-wide and per-source-key +handler gates across QUIC connections before reading a TCP target, plus a +finite request-header deadline. Separate +listener-wide TCP and UDP exit gates bound active target sockets across all +sessions. Accepted-but-unauthenticated HTTP/3 connections have a finite +authentication lifetime, and incoming unidirectional control streams have a +separate small per-connection cap. After QUIC Retry proves return-path +reachability, both global and source connection gates apply before handshake +state. The source key is an IPv4 address or IPv6 `/64`; clients sharing a NAT +or prefix intentionally share that budget. + +The TLS/TCP fallback has an independent listener-wide connection gate and a +per-source gate using the same IPv4-address / IPv6-`/64` keying. Both are held +from acceptance through handshake and for the complete relay lifetime. + +In `auto` mode, failure to establish or use the UDP session causes a new TCP +flow to be opened through one TLS 1.3 connection dedicated to that flow. A +valid relay-side destination error or authentication rejection is not retried +through fallback. Existing streams are never replayed or migrated: if UDP +fails after a stream has been handed to an application, that application must +retry and the new flow can use TLS. + +## UDP path + +SOCKS5 UDP ASSOCIATE is exposed only when the selected transport implements +datagrams. The client validates the UDP source against the SOCKS control +connection and rejects SOCKS fragmentation. Hysteria assigns a logical session +and fragments oversized Hysteria datagrams to the negotiated QUIC DATAGRAM +size. Before any per-session map or fragment slice is allocated, the relay +enforces global and per-source UDP-session caps across QUIC connections (IPv4 +address or IPv6 `/64`); +fragment count and reassembled payload size have fixed limits. +One logical Hysteria UDP payload is limited to 4,096 bytes. The SOCKS frontend +drops a larger payload before transport serialization without closing the UDP +association; ordinary Internet-MTU datagrams are unaffected. + +At the relay, every requested UDP destination is resolved and checked against +the same port, CIDR, private-address, and special-use policy as TCP. The +destination is resolved and filtered again for every outbound write, so DNS +changes cannot bypass the SSRF boundary. UDP sessions expire after a bounded +idle period. TLS/TCP fallback deliberately does not emulate UDP because doing +so would add cross-datagram head-of-line blocking and misleading semantics. + +## Congestion and loss recovery + +Congestion control chooses how quickly new data is put on the wire; loss +recovery decides when QUIC retransmits lost frames. They are separate: + +- BBR estimates delivered bandwidth and minimum RTT, derives a BDP, and paces + through STARTUP, DRAIN, PROBE_BW, and PROBE_RTT. +- Brutal paces at an explicitly configured rate and can compensate for the + measured ACK/loss rate. It is opt-in and is not congestion-fair. +- The QUIC transport declares loss using RFC 9002 packet and time thresholds + and sends PTO probes when ACK feedback stalls. AutoCAR does not replace this + with a proprietary predictor and does not add FEC. + +The controller runs in userspace on the AutoCAR QUIC connection. It does not +change the host's Linux `tcp_congestion_control`, accelerate unrelated sockets, +or act as a transparent TCP interception layer. + +## Security and probe behavior + +Both data paths use TLS 1.3. Clients must opt into either an explicit trust +anchor (`--ca`) or the operating-system roots (`--system-roots`), and normal +chain plus DNS/IP SAN verification is mandatory. A shared token is still +required and is compared in constant time; optional mTLS adds a client +certificate factor. + +Without Salamander, unauthenticated HTTP/3 requests receive a small neutral +cover page instead of an AutoCAR-specific error. The client also uses the +Hysteria transport's Chrome-oriented QUIC handshake fingerprint by default. +With Salamander, the UDP packet shape is obfuscated before it reaches QUIC; +ordinary HTTP/3 probes can no longer reach the cover site. Cover mode and +obfuscated-UDP mode are therefore alternative observable forms, not two layers +of one indistinguishable web service. + +These measures reduce obvious active-probe signatures; they do not make the +relay invisible. An observer can still see endpoints, timing, volume, packet +sizes, and TCP/UDP use, and can block or rate-limit the relay IP or all UDP. +AutoCAR makes no guarantee of being unidentifiable or unblockable. + +## Deliberate non-goals + +- no kernel-wide or transparent TCP acceleration; +- no proprietary Zeta-TCP prediction, redundant retransmission, or FEC; +- no BBRv2 or BBRv3 claim (the implemented model is BBRv1-derived); +- no port hopping, Mimic, or user-facing ECH configuration; +- no interception of destination HTTPS and no replacement for application + end-to-end encryption; and +- no promise that a relay improves every route or every workload. diff --git a/docs/BENCHMARK.md b/docs/BENCHMARK.md index 9ab8956..aa68ef5 100644 --- a/docs/BENCHMARK.md +++ b/docs/BENCHMARK.md @@ -1,118 +1,147 @@ # Benchmarking AutoCAR -AutoCAR includes a deterministic TCP source/sink so a direct route and the -dual-ended tunnel can be measured with the same payload. The benchmark reports -payload goodput in decimal Mbit/s. It is designed for repeatable comparisons, -not as proof that one transport is faster on every network. +AutoCAR includes a deterministic TCP source/sink so direct and relayed paths +can move the same payload. It reports payload goodput in decimal Mbit/s. The +tool is for repeatable comparisons; no single result proves that a relay or a +controller is faster on every network. ## Basic comparison -The benchmark server defaults to `127.0.0.1:9000`. For a same-host test: +The benchmark server defaults to `127.0.0.1:9000`: ```sh autocar bench-server ``` -A remote-path comparison needs an explicit non-loopback opt-in: +A remote-path comparison requires explicit non-loopback opt-in: ```sh autocar bench-server \ - --listen 0.0.0.0:9000 \ + --listen=0.0.0.0:9000 \ --allow-public-benchmark ``` `bench-server` has no authentication. A remote caller can make it send or -receive substantial traffic up to the configured limits. Use a host and cloud -firewall to allow port 9000 only from the intended client and relay addresses, -run it only for a controlled test window, and stop it immediately afterward. -The CLI defaults to at most 64 MiB per transfer and 16 concurrent transfers; -keep `--max-bytes` and `--max-connections` no higher than the experiment needs. +receive substantial traffic up to its limits. Restrict port 9000 to the test +client and relay with host/cloud firewalls, use the smallest practical +`--max-bytes` and `--max-connections`, and stop it after the test. -From the client host, measure the direct route: +Measure the direct route: ```sh autocar bench-client \ - --transport direct \ - --target bench.example.com:9000 \ - --mode download --bytes 8388608 --warmup 1 --iterations 7 --json + --transport=direct \ + --target=bench.example.com:9000 \ + --mode=download --bytes=8388608 \ + --warmup=1 --iterations=7 --json ``` -Then measure through an already running relay: +Measure the default Hysteria v2/BBR path through an already running relay: ```sh autocar bench-client \ - --transport quic \ - --server relay.example.com:443 \ - --ca relay-ca.crt \ - --token-file relay-token \ - --target bench.example.com:9000 \ - --mode download --bytes 8388608 --warmup 1 --iterations 7 --json + --transport=hy2 \ + --server=relay.example.com:443 \ + --ca=relay-ca.crt \ + --token-file=relay-token \ + --congestion=bbr --bbr-profile=standard \ + --upload-mbps=0 --download-mbps=0 \ + --target=bench.example.com:9000 \ + --mode=download --bytes=8388608 \ + --warmup=1 --iterations=7 --json ``` -Repeat with `--mode upload`. Use `--transport=tls` to characterize the TCP/TLS -fallback separately. For `--transport=auto`, the output transport label -describes the configured mode, not which path won an individual fallback -decision; use explicit modes when comparing transports. +`--transport=quic` is an alias for `hy2`. Use `--transport=tls` to characterize +the TCP/TLS fallback and `--transport=legacy-quic` only against a relay started +with `--quic-engine=legacy`. For `auto`, the JSON transport label describes the +configured mode rather than the path used by each individual flow; select an +explicit transport for performance comparisons. -The timer begins after the benchmark request header has been written and ends -after the payload plus a one-byte completion acknowledgement. Connection and -tunnel stream setup happen before that timer. For user-perceived latency, -measure the complete application operation separately. +Repeat with `--mode=upload`. The timer begins after the benchmark request +header is written and ends after the payload plus one-byte completion +acknowledgement. Connection and tunnel-stream setup occur before that timer. +Measure a complete real application operation separately when user-perceived +latency matters. + +## Controller matrix + +Do not compare only one controller on one path. A useful minimum matrix is: + +| Mode | Client flags | Relay flags | Question answered | +| --- | --- | --- | --- | +| Direct | `--transport=direct` | none | What does the unrelayed route deliver? | +| BBR conservative | `--congestion=bbr --bbr-profile=conservative`, bandwidths zero | matching BBR/profile, caps zero | Does a cautious model reduce queue/loss cost? | +| BBR standard | `--congestion=bbr --bbr-profile=standard`, bandwidths zero | matching BBR/profile, caps zero | Default model result | +| BBR aggressive | `--congestion=bbr --bbr-profile=aggressive`, bandwidths zero | matching BBR/profile, caps zero | Is extra startup pressure useful or harmful? | +| Brutal | truthful non-zero `--upload-mbps` and `--download-mbps` | `--allow-client-bandwidth` plus explicit non-zero negotiation ceilings | Does a reserved/controlled link benefit from a fixed negotiated rate? | +| Reno | `--congestion=reno`, bandwidths zero | `--congestion=reno`, caps zero | Loss-based baseline | +| TLS fallback | `--transport=tls` | TCP listener enabled | What is the reachability path's cost? | + +For Brutal, both relay caps and client measurements should be written into the +result metadata. An inflated capacity is not an optimization: it changes the +experiment into an unfair overload test. BBR profiles control the sender at +the endpoint where the flag is set, so record both endpoint configurations. + +For every row, exercise at least: + +- download and upload; +- short, medium, and bulk payloads; +- one flow and several concurrent flows; and +- clean, high-RTT, random-loss, burst-loss, and reordered profiles. ## Fair-test checklist -1. Pin the exact AutoCAR build, configuration, client, relay and benchmark - target for a comparison. -2. Keep the direct and tunneled destination identical. Document the different - physical routes and relay placement; a relay can improve routing, add a - detour, or both. -3. Run enough iterations in alternating order. Discard a declared number of - warmups and retain every raw result, not only the best value. -4. Test multiple payload sizes and concurrency levels. Short flows emphasize - setup and warm-state behavior; bulk transfers emphasize steady-state - congestion control. -5. Record RTT, loss, reordering, MTU, bandwidth, CPU utilization and time of - day. Confirm neither endpoint is CPU-limited. -6. Report median and the individual results. The emitted `p95_mbps` is the - 95th percentile of goodput, where larger is better; it is not a latency - percentile. -7. Repeat on the real production path. Emulation is useful for regression - testing but cannot reproduce every queue, middlebox or competing flow. - -Why QUIC can help: many flows reuse one authenticated connection and its -congestion state, and loss in one ordered QUIC stream does not impose -application-level head-of-line blocking on other streams. Why it may not help: -the relay adds processing and distance, a single large clean-path TCP flow can -already fill the link, and AutoCAR currently uses quic-go's default congestion -controller rather than claiming a custom BBR implementation. +1. Pin the AutoCAR commit, Go version, module versions, configuration, client, + relay, and benchmark target. +2. Keep the direct and tunneled destinations identical. Document both physical + routes and relay placement; a relay can improve routing or add a detour. +3. Alternate test order, declare warmups, run enough iterations, and retain + every raw result rather than only the best value. +4. Record RTT, random and burst loss, reordering, MTU, configured link rate, + CPU, memory, and time of day. Confirm neither endpoint is CPU-limited. +5. Report median plus all individual results. The emitted `p95_mbps` is the + 95th percentile of goodput, where larger is better; it is not latency p95. +6. Distinguish a warm shared QUIC connection from fresh direct TCP flows. That + is a real short-flow benefit, but it must be stated in the test description. +7. Repeat on the intended production path. Emulation catches regressions but + cannot reproduce every queue, middlebox, policer, or competing flow. + +Why Hysteria/QUIC can help: streams reuse a warm authenticated connection and +its BBR delivery/RTT model; pacing uses the inferred BDP; explicitly enabled +Fast Open can overlap the target response with initial writes; unrelated +streams avoid TCP-style cross-flow head-of-line blocking. Why it may not help: +the relay adds work and distance, the relay-to-destination leg is still a new +socket, and a clean direct TCP route may already fill the bottleneck. ## Reproducible Linux netem suite -The repository includes a root-only integration script. It creates isolated -client and relay network namespaces connected by a veth pair, applies the same -delay/loss/rate policy in both directions, and runs: +The root-only integration script creates isolated client and relay network +namespaces connected by a veth pair. Its current test matrix is: + +| Stage | Path profile | Cases | Pass condition | +| --- | --- | --- | --- | +| Bulk observation | 35 ms one-way delay on both interfaces, 0.5% independent loss each direction, 50 Mbit/s each direction | direct, Hysteria v2 (`quic` alias), TLS | every median is positive; ratios are retained | +| Controller gate | same lossy/rate-limited profile, repeated 4 MiB uploads and downloads | client-sender BBR/Reno/negotiated 15 Mbit/s Brutal; separate BBR and Reno relays for the relay sender | both upload and download BBR/Reno median ratios are at least 1.10; modes and negotiation are reported, and Brutal reaches at least 50% of its declared upload target | +| Cold fallback | same delay/rate, random loss removed, unused UDP port | `auto` Hysteria attempt followed by TLS | first command completes within finite deadlines | +| Short-flow acceleration gate | same delay/rate, loss-free, sequential 128 KiB downloads | fresh direct TCP vs warm Hysteria v2 connection | Hysteria median/direct median is at least 1.10 | +| Authentication | controlled namespace path | wrong CA and wrong token | both are rejected for the expected reason | +| Live UDP failure | first proxy request over Hysteria, then client UDP output is dropped | new TCP proxy flow in `auto` | new flow completes over TLS within the 10-second bound | +| Confidentiality smoke | pcap of Hysteria and fallback links | unique HTTP plaintext sentinel | sentinel is absent from both captures | -- direct, QUIC and TLS download measurements; -- an `auto` connection whose UDP address is initially unavailable, verifying - cold-start TCP/TLS fallback; -- an auto-mode proxy request that first succeeds over QUIC, followed by a - client-side UDP/7443 drop and a second bounded request over TCP/TLS; -- wrong-CA and wrong-token rejection checks; -- a controlled warm-QUIC short-flow acceleration profile; -- an HTTP proxy request containing a unique plaintext sentinel; and -- a packet capture assertion that the sentinel is absent from the client-relay - link. +The source and sink benchmark is TCP. SOCKS5 UDP ASSOCIATE, source validation, +datagram framing, and policy behavior are covered by Go integration tests; a +production UDP workload should also be measured with an application-specific +loss/jitter metric rather than TCP goodput. -On Linux with `iproute2`, `iptables`, `tcpdump`, `curl` and Python 3 installed: +Run the suite on Linux with `iproute2`, `iptables`, `tcpdump`, `curl`, Python 3, +and root privileges: ```sh make build sudo ./scripts/netem-integration.sh ./bin/autocar ``` -Defaults are 35 ms one-way delay on each side (approximately 70 ms base RTT), -0.5% independent loss in each direction, a 50 Mbit/s rate per direction, five -measured 1 MiB transfers and one warmup. They can be changed explicitly: +Override the declared profile explicitly: ```sh sudo env \ @@ -123,31 +152,35 @@ sudo env \ AUTOCAR_BENCH_ITERATIONS=9 \ AUTOCAR_BENCH_WARMUP=2 \ AUTOCAR_SHORT_FLOW_BYTES=131072 \ + AUTOCAR_SHORT_FLOW_ITERATIONS=9 \ + AUTOCAR_SHORT_FLOW_WARMUP=3 \ AUTOCAR_MIN_SHORT_FLOW_RATIO=1.10 \ + AUTOCAR_MIN_BBR_RENO_RATIO=1.10 \ + AUTOCAR_MIN_BRUTAL_TARGET_RATIO=0.50 \ AUTOCAR_ARTIFACT_DIR="$PWD/artifacts/netem" \ ./scripts/netem-integration.sh ./bin/autocar ``` -The script writes raw benchmark JSON, a comparison summary, process logs and -the pcap under `artifacts/netem`. The GitHub Actions netem workflow publishes -that directory as an artifact. - -The suite's general lossy-path bulk measurements are recorded without a speed -threshold. It separately applies one intentionally narrow acceptance profile: -loss-free high RTT, sequential 128 KiB downloads and three declared QUIC -warmups. Fresh direct TCP connections restart congestion state on every -iteration, while QUIC streams reuse the warm connection. The default gate -requires the warm QUIC median to be at least 1.10 times the direct median. This -demonstrates that the implemented connection-reuse acceleration mechanism is -effective under its stated conditions; it is not a universal production-speed -claim. `AUTOCAR_MIN_SHORT_FLOW_RATIO` can change the declared gate for a -different controlled environment, but a release should not lower it merely to -hide a regression. - -Performance on a shared virtual runner remains noisy. A broader release claim -should cite retained results from the intended path and configuration, not the -controlled CI profile alone. - -The pcap sentinel check is a useful regression smoke test, not a cryptographic -proof. The TLS 1.3 implementation, certificate validation and protocol threat -model remain the security basis. +The script writes raw JSON, a summary, process logs, and packet captures under +`artifacts/netem`. The GitHub Actions netem workflow publishes the directory +even when diagnosis is needed. + +## What the CI gate proves + +The generic bulk path measurements are observations rather than a universal +speed claim. Three controller-specific gates and one short-flow gate are narrow +and declared in advance: on the 0.5% lossy path, both client-side uploads and +relay-side downloads with BBR must beat their Reno baselines by at least 1.10, +and negotiated Brutal must deliver at least 50% of its truthful 15 Mbit/s +upload target; on the loss-free high-RTT path, sequential warm Hysteria +128 KiB downloads must beat fresh direct TCP by at least 1.10. These checks +demonstrate the selected mechanisms under those profiles only. + +Do not lower `AUTOCAR_MIN_SHORT_FLOW_RATIO` merely to hide a regression, and do +not publish the CI ratio as a universal production claim. BBR profile quality, +Brutal fairness, sustained high-loss behavior, and real-route improvement need +the broader retained matrix above. + +The pcap sentinel assertion is a regression smoke test, not a cryptographic +proof. TLS 1.3, verified X.509, token authentication, optional mTLS, and the +threat model remain the security basis. diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index 8a25d6f..71098bf 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -1,13 +1,13 @@ # Deployment guide -AutoCAR has two roles: a public relay close to the desired destinations and a -client-side process that exposes local SOCKS5 and HTTP(S) proxy listeners. The -preferred data path is QUIC over UDP; TCP with TLS 1.3 can listen on the same -port as a fallback. +AutoCAR has two roles: a relay near the desired destinations and a client-side +process exposing local SOCKS5 and HTTP(S) proxies. The default path is Hysteria +v2 over HTTP/3/QUIC on UDP. A separate TLS 1.3/TCP listener can use the same +numeric port for new-flow fallback. ## 1. Build and install -AutoCAR requires the Go version declared in `go.mod`. +AutoCAR requires the Go version declared in `go.mod`: ```sh make check @@ -15,53 +15,299 @@ make build sudo install -m 0755 bin/autocar /usr/local/bin/autocar ``` -The binary is statically buildable with `CGO_ENABLED=0`. Cross-platform -artifacts can be produced with `make cross-build`. Linux is the primary -deployment target; the local proxy also builds on macOS and Windows. +The binary is statically buildable with `CGO_ENABLED=0`; `make cross-build` +produces cross-platform artifacts. Linux is the primary relay target. The +local proxy also builds on macOS and Windows. -## 2. Create relay credentials +## 2. Create credentials -Create a dedicated account and a private configuration directory. Secret -files are rejected if they are accessible by group or other users. +Create a dedicated unprivileged account and a private configuration directory. +AutoCAR rejects secret/key files that are accessible by group or other users. ```sh sudo useradd --system --home /var/lib/autocar --shell /usr/sbin/nologin autocar sudo install -d -o autocar -g autocar -m 0700 /etc/autocar -sudo -u autocar autocar token --out /etc/autocar/relay-token + +sudo -u autocar autocar token --out=/etc/autocar/relay-token sudo -u autocar autocar cert \ - --hosts relay.example.com,203.0.113.10 \ - --cert /etc/autocar/server.crt \ - --key /etc/autocar/server.key + --hosts=relay.example.com,203.0.113.10 \ + --cert=/etc/autocar/server.crt \ + --key=/etc/autocar/server.key ``` -The bundled `cert` command creates an ECDSA P-256 self-signed certificate. Its -certificate file must be copied to each client over a trusted, independent -channel and used as `--ca`. For a public-CA certificate, clients may explicitly -select `--system-roots`; AutoCAR never offers an option to skip verification. +The bundled certificate command creates an ECDSA P-256 self-signed +certificate. Copy its certificate (never the private key) and the token to +each client over a trusted independent channel. Use the certificate as `--ca`. +For a public-CA certificate, clients can explicitly choose `--system-roots`. +There is no certificate-verification bypass. Default Chrome QUIC fingerprinting +supports ECDSA P-256/P-384 and RSA relay certificates but intentionally does not +advertise Ed25519. If an external issuer supplies an Ed25519 leaf, every +Hysteria client must use `--disable-chrome-parrot`; AutoCAR adds this exact hint +to a matching TLS handshake failure. Prefer P-256/P-384/RSA so the secure default can +remain enabled. -Keep the private key and token on the relay. Back them up as secrets, not in -the repository or container image. +Keep token, private key, optional client key, and obfuscation secret out of the +repository and container image. Give each secret a distinct random value; do +not reuse the relay token as a Salamander password or local-proxy password. -## 3. Run the relay +## 3. Run the relay with default BBR -For an initial foreground run on an unprivileged high port: +This is a complete foreground example using the default Hysteria engine, +BBRv1 `standard`, UDP proxying, neutral HTTP/3 cover, and TLS fallback: ```sh sudo -u autocar autocar server \ - --listen :8443 \ - --tcp-listen :8443 \ - --cert /etc/autocar/server.crt \ - --key /etc/autocar/server.key \ - --token-file /etc/autocar/relay-token + --listen=:8443 \ + --quic-engine=hy2 \ + --tcp-listen=:8443 \ + --cert=/etc/autocar/server.crt \ + --key=/etc/autocar/server.key \ + --token-file=/etc/autocar/relay-token \ + --congestion=bbr \ + --bbr-profile=standard \ + --max-upload-mbps=0 \ + --max-download-mbps=0 \ + --max-connections=256 \ + --max-client-connections=32 \ + --max-client-fallback-connections=32 \ + --max-streams=1024 \ + --max-uni-streams=8 \ + --max-outbound-tcp=1024 \ + --max-outbound-udp=256 \ + --max-client-tcp-handlers=128 \ + --max-client-udp-sessions=64 \ + --handshake-timeout=10s \ + --udp-idle-timeout=60s \ + --masquerade-name="Example Service" +``` + +The relay ignores client bandwidth hints by default, so this selects BBR rather +than a guessed or client-forced Brutal rate. Enabling client hints is a separate +explicit opt-in described below and requires two finite relay ceilings. +`--bbr-profile=conservative` is a good first choice on shared/shallow-buffer +links; `aggressive` should be reserved for measured controlled paths. +`--congestion=reno` offers a loss-based baseline. Controller settings affect +only traffic sent by the endpoint on which they are configured. + +Production deployments often use port 443. Permit both UDP and TCP. Binding a +low port as non-root requires a narrow capability such as +`CAP_NET_BIND_SERVICE`, or one-to-one UDP and TCP port forwarding. A load +balancer must use UDP/TCP pass-through; terminating TLS or HTTP/3 in front of +AutoCAR changes the required end-to-end protocol. + +### Destination policy + +The relay rejects loopback, link-local, multicast, unspecified, RFC1918, ULA, +CGNAT, translation, documentation, benchmarking, reserved, and other +special-use destinations by default. It also denies ports 25, 465, and 587. + +- `--allow-private` permits RFC1918/ULA/CGNAT but never loopback, link-local, or + the built-in special-use denylist. Use it only on a relay dedicated to + trusted users. +- `--deny-cidrs` adds deployment-specific IPs/CIDRs such as cloud control-plane + ranges. +- `--deny-ports=none` clears the port denylist and is a deliberate security + policy change. +- `--disable-udp` disables Hysteria UDP sessions and SOCKS5 UDP ASSOCIATE. + +TCP hostnames are resolved once into approved numeric dial candidates. UDP is +resolved and filtered during the permission check and again on every outbound +write, preventing DNS rebinding from bypassing the egress policy. + +## 4. Run the client + +The default `auto` mode first uses Hysteria v2 and falls back to TLS/TCP for new +TCP flows when UDP is unavailable: + +```sh +autocar client \ + --server=relay.example.com:443 \ + --transport=auto \ + --ca=/etc/autocar/relay-ca.crt \ + --token-file=/etc/autocar/relay-token \ + --congestion=bbr \ + --bbr-profile=standard \ + --upload-mbps=0 \ + --download-mbps=0 \ + --fast-open=false \ + --max-pending-opens=256 \ + --socks=127.0.0.1:1080 \ + --http=127.0.0.1:8080 +``` + +`--transport=hy2` and its alias `--transport=quic` require UDP and do not fall +back. `--transport=tls` diagnoses the fallback directly. If TCP uses a +different address, set `--fallback-server`. + +Fast Open defaults to off so a destination refusal is returned before the +application sees an established proxy connection. `--fast-open=true` can +reduce a round trip for write-first protocols, but defers the relay dial result +until the first read and should be enabled only after testing application error +handling. + +`--quic-attempt-timeout` bounds the primary UDP attempt, +`--fallback-cooldown` controls how long new TCP flows prefer TLS after a UDP +failure, and `--open-timeout` bounds the whole proxy open. In `auto`, keep +`0 < --quic-attempt-timeout < --open-timeout` so fallback retains time. A +relay destination error or authentication failure does not open the fallback +circuit. `--max-pending-opens` bounds Hysteria core operations that cannot be +interrupted through its public API; canceled callers return immediately, but +their late workers retain a slot until the core returns or the session closes. + +Fallback is not live migration. A stream already returned to an application +fails if its QUIC connection becomes unusable; the application must retry, and +the new TCP flow can use TLS. SOCKS5 UDP has no TLS/TCP fallback. + +Application examples: + +```sh +curl --socks5-hostname 127.0.0.1:1080 https://example.com/ +curl --proxy http://127.0.0.1:8080 https://example.com/ +``` + +The SOCKS5 frontend supports CONNECT and UDP ASSOCIATE, but not BIND. The HTTP +proxy supports absolute-form `http://` and CONNECT. Absolute-form `https://` +is rejected; HTTPS destinations use CONNECT, leaving application TLS end to +end. AutoCAR never installs an interception CA. + +## 5. Opt into negotiated Brutal + +Use Brutal only after measuring the access link and only where reserving that +rate is permitted. The client values request the two directional rates; relay +values set negotiation ceilings for cooperating clients: + +```sh +# Relay: bound negotiated Brutal targets +autocar server [relay options] \ + --allow-client-bandwidth \ + --max-upload-mbps=20 \ + --max-download-mbps=100 + +# Client: truthful measured capacities +autocar client [client options] \ + --upload-mbps=20 \ + --download-mbps=100 +``` + +A zero client value keeps BBR/Reno in that direction. With non-zero client and +relay values, the lower value wins. Set `--disable-loss-compensation` on both +roles to disable the five-second ACK/loss compensation in both directions. +Without `--allow-client-bandwidth`, the server ignores all client hints and +forces its configured BBR/Reno behavior. The opt-in is rejected unless both +server ceilings are non-zero. + +Brutal can send about 1.25 times the requested rate under measured loss. It is +not congestion-fair; an exaggerated value can starve other traffic, waste +capacity, and trigger a provider policer. Relay negotiation ceilings are an +important guardrail for the official client, but they are not a traffic +policer: use host or cloud rate limiting when a hard, non-bypassable ceiling is +required. + +## 6. Choose one probe-resistance form + +### Plain HTTP/3 cover + +With no obfuscation secret, the relay is a valid HTTP/3 endpoint. +Unauthenticated probes receive the neutral site named by `--masquerade-name`. +The client uses Hysteria's Chrome-oriented QUIC fingerprint by default; +normally `--disable-chrome-parrot` is intended for diagnostics. It is also +required when the relay deliberately uses an Ed25519 certificate; ECDSA P-256 +(including `autocar cert`) and RSA work with the default fingerprint. + +### Salamander-obfuscated UDP + +Generate a separate secret and install the same 0600 file on both endpoints: + +```sh +sudo -u autocar autocar token --out=/etc/autocar/obfs-password + +# Add to both relay and client commands: +--obfs-password-file=/etc/autocar/obfs-password ``` -Permit both UDP and TCP on the chosen port. Production deployments commonly -use 443 because restrictive networks are more likely to permit it. Binding a -port below 1024 as a non-root process requires a narrowly scoped capability -such as `CAP_NET_BIND_SERVICE`, or a firewall/load-balancer mapping from 443 to -8443. +The alternative `AUTOCAR_OBFS_PASSWORD` environment variable is supported but +must be protected from service-manager and process-environment disclosure. +Salamander requires the `hy2` engine and changes the outer UDP packet form. +An ordinary HTTP/3 probe can no longer reach the inner cover page; cover and +Salamander are alternative observable modes. The TCP fallback is unaffected. + +Neither mode guarantees that a network operator cannot classify, rate-limit, +or block the relay. Endpoint IP, timing, volume, packet sizes, and TCP/UDP use +remain visible. -A minimal hardened systemd service is: +## 7. Firewall and UDP rate limits + +Allow both protocols on the relay's selected port, restricted by source ranges +when possible: + +```sh +sudo ufw allow from 198.51.100.0/24 to any port 443 proto udp +sudo ufw allow from 198.51.100.0/24 to any port 443 proto tcp +``` + +`--max-connections` bounds all accepted QUIC sessions, including +unauthenticated cover traffic; `--max-client-connections` prevents one +return-path-validated source key from occupying the whole connection budget. +The key is one IPv4 address or one IPv6 `/64`, preventing ordinary IPv6 address +rotation from bypassing the budget. +`--max-streams` bounds incoming streams per QUIC connection. +The same value is the global TLS/TCP fallback connection gate, while +`--max-client-fallback-connections` prevents one IPv4 address or IPv6 `/64` +from occupying it during the TLS handshake or relay lifetime. Its zero default +selects the smaller of 32 and `--max-streams`; set a positive value only when +the measured client concurrency requires a different source budget. +`--max-uni-streams` separately limits Hysteria/HTTP/3 control +streams, while `--handshake-timeout` closes an accepted connection that does +not authenticate in time and also bounds the initial TCP target header. A +global pending-TCP-handler gate prevents slow streams from bypassing the +listener-wide `--max-outbound-tcp` backstop. The +`--max-client-tcp-handlers` source budget is shared across QUIC connections, +so one authenticated source cannot occupy the complete global gate. HTTP +request headers are rejected above 16 KiB before body allocation. UDP sessions +are admitted before fragment state is allocated, with both `--max-outbound-udp` global and +`--max-client-udp-sessions` per-authenticated-source-key limits shared across +QUIC connections. IPv4 is keyed per address and IPv6 per `/64`. Clients behind +the same NAT or routed prefix share all three per-source +connection/TCP/UDP budgets; +increase them only after measuring legitimate concurrency, while retaining the +global caps as protection against source-address rotation. +Fragment count and reassembled size are also bounded. Capacity is released on +close. QUIC Retry validates the source address before a bounded handshake slot +is allocated. Initial packets still consume kernel and link work, so combine +these controls with a host/cloud UDP rate guard. The following nftables fragment is a template +for an existing firewall. It drops per-source UDP above 25 MiB/s (about +210 Mbit/s) with an 8 MiB burst, then leaves the accept/drop policy to the +site's normal filter chain: + +```nft +table inet autocar_guard { + chain input { + type filter hook input priority -5; policy accept; + + udp dport 443 meter autocar_udp4 { + ip saddr limit rate over 25 mbytes/second burst 8 mbytes + } drop + + udp dport 443 meter autocar_udp6 { + ip6 saddr limit rate over 25 mbytes/second burst 8 mbytes + } drop + } +} +``` + +Validate nftables syntax on the target distribution before loading it. Set the +limit above the largest authorized Brutal sender rate plus protocol overhead; +a lower firewall ceiling silently invalidates the negotiated rate. Also apply +provider edge limits because host rules cannot recover bandwidth already +consumed upstream. Keep SSH/management access in a separately tested rule set. + +SOCKS5 UDP ASSOCIATE returns a dynamically allocated local UDP port. If the +client proxy runs in a container, ordinary fixed TCP port publishing does not +publish that dynamic UDP endpoint. Use a host-local client process or a +carefully firewalled host-network deployment for applications that need SOCKS +UDP. + +## 8. Hardened systemd relay ```ini [Unit] @@ -72,7 +318,7 @@ Wants=network-online.target [Service] User=autocar Group=autocar -ExecStart=/usr/local/bin/autocar server --listen=:443 --tcp-listen=:443 --cert=/etc/autocar/server.crt --key=/etc/autocar/server.key --token-file=/etc/autocar/relay-token +ExecStart=/usr/local/bin/autocar server --listen=:443 --quic-engine=hy2 --tcp-listen=:443 --cert=/etc/autocar/server.crt --key=/etc/autocar/server.key --token-file=/etc/autocar/relay-token --congestion=bbr --bbr-profile=standard --max-upload-mbps=0 --max-download-mbps=0 --masquerade-name=Example-Service Restart=on-failure RestartSec=3 AmbientCapabilities=CAP_NET_BIND_SERVICE @@ -92,168 +338,108 @@ MemoryDenyWriteExecute=true WantedBy=multi-user.target ``` -The relay rejects loopback, link-local, multicast, unspecified, RFC1918, ULA, -CGNAT and IANA special-use destinations by default. It also denies ports 25, -465 and 587. `--allow-private` permits private/ULA/CGNAT targets but never -loopback, link-local or the explicitly denied special-use prefixes; use it only -for a relay dedicated to trusted users. `--deny-cidrs` adds deployment-specific -blocked IPs or CIDRs (for example, a cloud-provider control-plane range). -`--deny-ports=none` removes the port denylist and should be treated as a -deliberate security-policy change. - -## 4. Run the client +Run `systemd-analyze security autocar.service`, adapt restrictions to the host, +then test UDP and TCP separately. A watchdog should probe both because a green +TCP fallback does not prove that Hysteria UDP is reachable. -Place the relay certificate and the same token on the client, with the token -owned by the local AutoCAR account and mode `0600`: +## 9. Local proxy exposure -```sh -autocar client \ - --server relay.example.com:443 \ - --transport auto \ - --ca /etc/autocar/relay-ca.crt \ - --token-file /etc/autocar/relay-token \ - --socks 127.0.0.1:1080 \ - --http 127.0.0.1:8080 -``` - -`auto` first uses QUIC. If that path fails, it uses the TLS/TCP listener; a -short circuit-breaker cooldown prevents every new flow from repeatedly -waiting for an unavailable UDP path. `--transport=quic` and -`--transport=tls` are useful for diagnosis. If the TCP fallback uses a -different address, set `--fallback-server`. `--quic-attempt-timeout` bounds -the whole QUIC phase, while `--dial-timeout` bounds network establishment and -`--open-timeout` bounds the overall proxy open operation. Keep the overall -timeout comfortably larger than the QUIC budget so TLS has time to complete; -the defaults also leave the relay's destination dial timeout inside that -budget on a normally responsive path. - -Fallback is connection-establishment behavior, not live stream migration. If -UDP disappears after a QUIC stream has already been returned to an -application, that stream fails and the application must retry; newly opened -streams use TLS while the QUIC circuit is open. Arbitrary TCP bytes cannot be -safely replayed onto a different transport without application cooperation. - -Application examples: +Unauthenticated local proxy listeners are restricted to loopback. To expose a +listener on another interface, configure credentials and explicitly +acknowledge that SOCKS username/password and HTTP Basic are cleartext on those +listeners: ```sh -curl --socks5-hostname 127.0.0.1:1080 https://example.com/ -curl --proxy http://127.0.0.1:8080 https://example.com/ +AUTOCAR_PROXY_USER=alice autocar client [client options] \ + --proxy-password-file=/etc/autocar/proxy-password \ + --allow-public-plaintext \ + --socks=0.0.0.0:1080 \ + --http=0.0.0.0:8080 ``` -The SOCKS5 frontend currently supports CONNECT, not BIND or UDP ASSOCIATE. -The HTTP proxy supports absolute-form `http://` requests and CONNECT. -Absolute-form `https://` is rejected: HTTPS destinations must use CONNECT, so -their application TLS remains end to end between the application and -destination. AutoCAR does not install a CA or intercept destination TLS. - -The optional `--https` listener encrypts the hop from an application to the -local proxy. It requires `--proxy-cert` and `--proxy-key`: +Prefer the TLS-protected local HTTPS proxy across an untrusted LAN: ```sh -autocar cert --hosts localhost,127.0.0.1 --cert proxy.crt --key proxy.key -autocar client [relay options] \ - --socks= --http= --https 127.0.0.1:8444 \ - --proxy-cert proxy.crt --proxy-key proxy.key -curl --proxy https://127.0.0.1:8444 --proxy-cacert proxy.crt https://example.com/ -``` +autocar cert \ + --hosts=localhost,127.0.0.1 \ + --cert=proxy.crt --key=proxy.key -Unauthenticated proxy listeners are restricted to loopback. To expose one on -another interface, set a username and a permission-restricted password file: +autocar client [client options] \ + --socks= --http= \ + --https=127.0.0.1:8444 \ + --proxy-cert=proxy.crt \ + --proxy-key=proxy.key -```sh -AUTOCAR_PROXY_USER=alice autocar client [relay options] \ - --proxy-password-file /etc/autocar/proxy-password \ - --allow-public-plaintext \ - --socks 0.0.0.0:1080 --http 0.0.0.0:8080 +curl --proxy https://127.0.0.1:8444 \ + --proxy-cacert proxy.crt https://example.com/ ``` -SOCKS5 uses username/password authentication and the HTTP proxy uses Basic -proxy authentication. Both transmit local-proxy credentials without transport -encryption, which is why a non-loopback listener requires the explicit -`--allow-public-plaintext` acknowledgement. These credentials protect the -local listener; they do not replace the relay token or TLS certificate -validation. Add host firewall rules even when authentication is enabled, and -prefer the HTTPS proxy listener across any untrusted local network. +Local proxy credentials protect the listener; they do not replace the relay +token or X.509 verification. Use host firewall rules even with authentication. -## 5. Optional mutual TLS +## 10. Optional mutual TLS -The shared token is mandatory. mTLS adds a client-certificate factor. Generate -a separate client certificate, install its certificate (not its key) as the -relay client trust anchor, and start the roles with: +The token remains mandatory. mTLS adds a client-certificate factor: ```sh # Relay -autocar server [server options] --client-ca /etc/autocar/client.crt +autocar server [relay options] \ + --client-ca=/etc/autocar/client-ca.crt # Client autocar client [client options] \ - --client-cert /etc/autocar/client.crt \ - --client-key /etc/autocar/client.key + --client-cert=/etc/autocar/client.crt \ + --client-key=/etc/autocar/client.key ``` -For multiple clients, use a conventional private CA and issue distinct client -certificates so identities remain auditable and can be migrated independently. -Protocol v1 does not implement CRL, OCSP, or a certificate denylist; revoking a -compromised client therefore requires rotating the accepted client CA and the -remaining client certificates (and rotating the shared token when exposed). - -## 6. Containers and Compose +Issue distinct client certificates from a private CA so identities can be +audited and rotated independently. AutoCAR has no CRL, OCSP, or certificate +denylist; revocation requires rotating the accepted client CA/certificates and +the token if it was exposed. -The image uses a multi-stage build and a `scratch` runtime. It includes the -system CA bundle, contains only the binary and CA file, and runs as numeric UID -and GID 65532 with no Linux capabilities. +## 11. Wire migration from original AutoCAR QUIC -`docker-compose.yml` provides separate `server` and `client` profiles. Prepare -the expected files under `./secrets`, keep token/password/key files at mode -`0600`, and make them readable by the configured container UID. For example: +`--transport=quic` now aliases Hysteria v2. It does not speak the original +AutoCAR v1 QUIC wire. Upgrade UDP client and server together. For temporary +compatibility, make both selections explicit: ```sh -mkdir -p secrets -bin/autocar token --out secrets/relay-token -bin/autocar cert --hosts relay.example.com \ - --cert secrets/server.crt --key secrets/server.key -cp secrets/server.crt secrets/relay-ca.crt -bin/autocar token --out secrets/proxy-password -sudo chown -R 65532:65532 secrets - -docker compose --profile server up --build -d +# Old v1 UDP engine +autocar server [relay options] --quic-engine=legacy +autocar client [client options] --transport=legacy-quic ``` -For a client host, set `AUTOCAR_RELAY`, `AUTOCAR_SERVER_NAME` and a non-default -`AUTOCAR_PROXY_USER`, then run `docker compose --profile client up --build -d`. -Compose publishes the local proxy ports only on host loopback, while proxy -authentication is still mandatory inside the container because its listener -binds the container interface. The Compose command also sets -`--allow-public-plaintext` explicitly: SOCKS5 username/password and HTTP Basic -proxy credentials are not encrypted on that container-side listener. The -loopback-only host publishing and Docker network boundary are therefore part -of this example's security model. Do not change those port mappings to a -public host address; use the local HTTPS proxy listener or another encrypted -hop if clients must cross an untrusted network. - -If host files cannot be owned by UID 65532, set `AUTOCAR_UID` and -`AUTOCAR_GID` to their owner. The Dockerfile's default process remains -non-root; do not set the Compose user to root merely to work around secret-file -permissions. - -## 7. Operations and limits - -- Rotate a token by updating both ends during a coordinated restart. There is - no multi-token grace period in protocol v1. -- The current release emits lifecycle and fatal-command logs, but deliberately - has no unauthenticated metrics endpoint and does not log every hostile - request. Monitor restarts with the process supervisor, use firewall counters - and bounded external probes for reachability, and alert on host CPU, memory, - file-descriptor and network saturation. Built-in aggregate auth/dial/stream - metrics remain future work; logs intentionally avoid payload, credentials, - destinations and internal dial details. -- Keep Go and module dependencies patched. CI runs unit/race tests and CodeQL; - those checks complement rather than replace dependency and host patching. -- Test UDP and TCP reachability independently after every firewall, NAT or - load-balancer change. -- A relay sees requested destinations and any destination-side plaintext. - Continue using HTTPS, SSH or another end-to-end protocol for sensitive data. -- Observers still see relay IPs, packet sizes, timing and whether UDP/TLS is in - use. No protocol can promise that a network operator will never rate-limit - or block it. AutoCAR's TCP/TLS fallback improves reachability but is not an - undetectability guarantee. +One address cannot host both UDP engines. For a staged migration, bind Hysteria +to a second UDP port, move clients, then retire the legacy port. The separate +TLS/TCP fallback remains AutoCAR v1, so `auto` can still provide new-flow TCP +reachability during a UDP mismatch. Legacy QUIC has no BBR/Brutal integration +or UDP ASSOCIATE. + +## 12. Containers and operations + +The image uses a multi-stage build, a `scratch` runtime, numeric UID/GID 65532, +and no Linux capabilities. `docker-compose.yml` publishes relay UDP and TCP and +publishes local proxy TCP listeners only on host loopback. Prepare 0600 files +under `./secrets`, make them readable by the configured container UID, and use +the `server` or `client` profile. Do not publish the cleartext client proxy on +a public host address. + +Operational checklist: + +- rotate relay token and optional obfuscation secret with a coordinated restart; +- independently test Hysteria UDP and fallback TCP after every firewall, NAT, + certificate, or load-balancer change; +- retain the controller/capacity configuration with benchmark results and + remeasure after route or provider changes; +- monitor process restarts, CPU, memory, file descriptors, UDP drops, firewall + counters, and link saturation; +- keep Go, modules, the host kernel, and container base/build images patched; + and +- continue using HTTPS, SSH, or another end-to-end application protocol because + the relay necessarily sees requested destinations and destination-side + plaintext. + +AutoCAR currently has no unauthenticated metrics endpoint and intentionally +does not log payloads or credentials. It also has no port hopping, Mimic, +user-facing ECH configuration, kernel-transparent TCP mode, or FEC. Do not +describe it as unidentifiable or unblockable. diff --git a/docs/PROTOCOL.md b/docs/PROTOCOL.md index 8f2eb2b..10c8950 100644 --- a/docs/PROTOCOL.md +++ b/docs/PROTOCOL.md @@ -1,30 +1,160 @@ -# AutoCAR wire protocol v1 +# AutoCAR transports and wire protocols -This document describes the protocol implemented by `internal/protocol`. It is -intended to make compatibility and security review possible; it is not a -promise that every internal Go API is stable. +AutoCAR currently has two UDP wire protocols plus one TCP fallback protocol. +The CLI transport name and the bytes on the wire must not be confused. -## Transport binding +## Compatibility matrix + +| Client selection | Server selection | Carrier | Wire protocol | Datagram proxying | +| --- | --- | --- | --- | --- | +| `auto` (default) | `--quic-engine=hy2` (default) | UDP first, TCP fallback | Hysteria v2; AutoCAR v1 on fallback | Yes on the UDP path | +| `hy2` or `quic` | `--quic-engine=hy2` | UDP | Hysteria v2 | Yes | +| `legacy-quic` | `--quic-engine=legacy` | UDP | AutoCAR v1 | No | +| `tls` | either UDP engine | TCP | AutoCAR v1 | No | + +`quic` is an alias for `hy2`. It exists for CLI continuity, not wire +compatibility with the original AutoCAR QUIC engine. A client and server that +select different UDP engines cannot complete the UDP handshake. + +## Default Hysteria v2 wire + +The default engine delegates its wire format and state machine to +[`github.com/apernet/hysteria/core/v2` v2.12.1](https://github.com/apernet/hysteria/tree/14e9fff1d972ab0187ac7fcf75b9514dc8664065/core). +It is Hysteria v2 over HTTP/3 and QUIC, rather than an AutoCAR-specific framing +layer. This section records AutoCAR's binding and security policy; the upstream +[Hysteria v2 protocol documentation](https://github.com/apernet/hysteria/blob/14e9fff1d972ab0187ac7fcf75b9514dc8664065/PROTOCOL.md) +is the interoperability reference. + +### Authentication and transport negotiation + +The client establishes HTTP/3 with TLS 1.3 and sends the Hysteria authentication +request containing the shared token and the client's receive-rate hint. On +success, the relay returns whether UDP is enabled and its receive-rate +advertisement. The two advertisements independently select the client-to-relay +and relay-to-client sender rate. + +For each direction: + +1. by default the relay ignores all client capacity hints, keeps the configured + model controller (BBR by default, or Reno), and reports automatic bandwidth + selection; +2. the operator must explicitly set `--allow-client-bandwidth` and finite + non-zero upload/download ceilings before the relay accepts Brutal hints; and +3. in that opt-in mode, a non-zero client hint selects Brutal at the lower of + the client value and corresponding server ceiling, while a zero hint keeps + BBR/Reno for that direction. + +The rate values on the Hysteria wire are bytes per second. AutoCAR's CLI accepts +decimal Mbit/s and converts with `Mbit/s * 1,000,000 / 8`. + +| Hysteria field | AutoCAR CLI meaning | +| --- | --- | +| Client `MaxTx` | `--upload-mbps`, client to relay | +| Client `MaxRx` | `--download-mbps`, relay to client | +| Server `MaxTx` | `--max-download-mbps`, relay-to-client Brutal negotiation ceiling | +| Server `MaxRx` | `--max-upload-mbps`, client-to-relay Brutal negotiation ceiling | + +These server values constrain negotiation with the official client; they are +not a packet policer and do not impose a hard throughput limit on a modified +or zero-hint client. Enforce non-bypassable limits in the host or cloud network. + +Unknown, malformed, and unauthenticated HTTP/3 requests are passed to the +configured cover handler. AutoCAR's built-in handler returns a neutral page for +`GET /` and `HEAD /`, with ordinary not-found responses elsewhere. It does not +reveal whether a supplied token was close to valid. + +### TCP streams and Fast Open + +Each proxied TCP flow uses a Hysteria bidirectional QUIC stream. Its bounded +request identifies the destination, and the relay returns success or a generic +dial error before opaque byte forwarding. Fast Open is disabled by default. +With client `--fast-open=true`, the connection is returned after the target +request is written; application writes may proceed while the response is +deferred until the first read. Fast Open does not bypass TLS, token +authentication, outbound policy, or the destination dial. + +An authenticated relay error is an application result, not evidence that UDP +is unavailable. `auto` therefore does not duplicate that request through TLS. + +### UDP datagrams + +SOCKS5 UDP ASSOCIATE creates one Hysteria logical UDP session. Each QUIC +DATAGRAM carries a session identifier, packet/fragment identifiers, target +address, and payload. Hysteria fragments messages that exceed the available +QUIC DATAGRAM payload and reassembles them inside the same logical session. +This fragmentation is bounded: one logical payload may be at most 4,096 bytes. +The SOCKS frontend drops a larger payload before calling the Hysteria transport +and keeps the association alive. DATAGRAM delivery is intentionally unreliable +and unordered; QUIC does not +retransmit a lost UDP payload. + +The SOCKS request's `DST.ADDR` and `DST.PORT` describe the expected local UDP +source as specified by RFC 1928. An unspecified address is bound to the TCP +control peer; a concrete IP must equal that peer, and a domain must resolve to +it. A non-zero port is enforced. Port zero remains the normal dynamic form and +is locked to the first valid datagram. The actual packet source is checked even +after request validation, so DNS cannot authorize a different sender. + +The relay bounds session/fragment state before allocation, checks the first +complete target before creating an outbound socket, and re-resolves plus +revalidates the destination on every outbound datagram. Replies are accepted +only from numeric destinations that the session previously approved and wrote +to successfully. This per-session destination set is capped at 256 entries and +never evicts an entry: after it fills, existing destinations remain usable but +new destinations fail closed. A failed socket write does not authorize its +source address for replies. UDP is unavailable when the relay uses +`--disable-udp`, when the client +selects a v1 transport, or when `auto` has only its TCP fallback path available. + +### Loss signals are standard QUIC + +Hysteria's congestion controller consumes ACKed/lost packet events from the +QUIC transport. Loss declaration itself follows RFC 9002: a packet-number +threshold of three, a time threshold of 9/8 of the relevant RTT estimate, and +PTO probes with exponential backoff. These mechanisms provide rapid standard +loss recovery; they are not advertised as Zeta-TCP's proprietary prediction or +reverse-control implementation. + +### Cover and Salamander forms + +In plain mode, the UDP endpoint is valid HTTP/3 and unauthenticated probes can +receive the cover site; Hysteria's Chrome-oriented QUIC handshake fingerprint +is enabled by default. Its signature list supports ECDSA P-256/P-384/RSA relay +certificates, not Ed25519. `--disable-chrome-parrot` exists for diagnostics and +is required for an Ed25519 relay certificate; X.509 verification remains +mandatory either way. + +When both endpoints load the same `--obfs-password-file` (or +`AUTOCAR_OBFS_PASSWORD`), Salamander wraps the UDP packet connection. A normal +HTTP/3 client then cannot reach the inner cover handler. This is an alternate +obfuscated packet form, not HTTP/3 masquerading and obfuscation simultaneously +visible on the network. Salamander does not affect the separate TLS/TCP +fallback. + +## AutoCAR wire protocol v1 + +Protocol v1 remains the TLS/TCP fallback format and the explicit legacy QUIC +format. It is implemented by `internal/protocol` and retained for controlled +migration and fallback, not used by default Hysteria UDP. + +### Transport binding Protocol v1 is carried over either: -- one bidirectional stream in a QUIC v1 or v2 connection, as negotiated by - the pinned quic-go transport, or +- one bidirectional stream in a legacy QUIC v1 or v2 connection, or - one TLS-over-TCP connection dedicated to a single proxied stream. -Both transports require TLS 1.3 and negotiate the ALPN value `autocar/1`. -Normal X.509 chain and DNS/IP SAN verification is mandatory on the client. -The server can additionally require an mTLS client certificate. QUIC 0-RTT -and QUIC DATAGRAM are disabled, so a CONNECT request is not sent as replayable -early data. +Both require TLS 1.3 and negotiate ALPN `autocar/1`. Normal X.509 chain and +DNS/IP SAN verification is mandatory. The server can additionally require an +mTLS client certificate. Legacy QUIC disables 0-RTT and QUIC DATAGRAM. -Each QUIC stream or fallback TLS connection contains exactly one CONNECT -exchange followed by an unframed TCP byte stream. Integers are unsigned and -encoded in network byte order (big endian). +Each stream or TLS connection contains one CONNECT exchange followed by an +unframed TCP byte stream. Integers are unsigned and encoded in network byte +order (big endian). -## Common header +### Common header -Every request and response starts with this 12-byte header: +Every v1 request and response starts with this 12-byte header: | Offset | Size | Field | Value | | ---: | ---: | --- | --- | @@ -35,42 +165,36 @@ Every request and response starts with this 12-byte header: | 8 | 2 | Length 1 | Kind-specific body length | | 10 | 2 | Length 2 | Kind-specific body length or zero | -Readers reject a bad magic value, unsupported version or unexpected kind. -Every length is checked against its semantic limit before allocation. - -## CONNECT request +Readers reject a bad magic value, unsupported version, unexpected kind, or a +length beyond its semantic limit before allocation. -A request uses kind `1` and this layout: +### CONNECT request | Header field | Meaning | | --- | --- | | Flags | Network: `1` = `tcp`, `2` = `tcp4`, `3` = `tcp6` | -| Length 1 | Shared-token byte length, from 16 through 1024 | -| Length 2 | destination byte length, from 1 through 1024 | -| Body | token bytes followed immediately by destination bytes | - -The destination is a Go `net.SplitHostPort`-compatible `host:port` value. -IPv6 literals therefore use brackets, for example `[2001:db8::1]:443`. NUL -bytes are forbidden. Protocol framing does not separately validate UTF-8. A -hostname is preserved for resolution by the relay. +| Length 1 | Shared-token byte length, 16 through 1024 | +| Length 2 | Destination byte length, 1 through 1024 | +| Body | Token bytes followed immediately by destination bytes | -The token is transported only after the encrypted channel has been -established. The relay compares a SHA-256 digest of the presented token with -the configured token digest using a constant-time comparison. The token is an -authorization factor, not a replacement for certificate verification. +The destination is a `host:port` value compatible with Go's +`net.SplitHostPort`. IPv6 literals use brackets, for example +`[2001:db8::1]:443`. NUL bytes are forbidden. Hostnames are preserved for +relay-side resolution. -## Response +The token is sent only inside the encrypted channel. The relay compares a +SHA-256 digest of the presented token with the configured digest using a +constant-time comparison. It is an authorization factor, not a replacement +for certificate verification. -A response uses kind `2` and this layout: +### Response | Header field | Meaning | | --- | --- | | Flags | Status code | | Length 1 | Optional human-readable message length, 0 through 1024 | | Length 2 | Reserved; must be zero | -| Body | message bytes | - -Status values are: +| Body | Message bytes | | Value | Name | Meaning | | ---: | --- | --- | @@ -81,28 +205,37 @@ Status values are: | 4 | Busy | Concurrent-stream capacity was exhausted | | 5 | Internal | Reserved for a relay-side internal failure | -Relay errors deliberately avoid returning resolver, host-topology or operating -system details. Clients treat a valid non-OK response as a remote error, not a -transport outage; it therefore does not trigger a retry through the TCP/TLS -fallback. +Relay errors avoid resolver, host-topology, and operating-system details. A +valid non-OK response is a remote error and does not trigger another transport. + +### Relay phase and shutdown + +After OK, both sides clear the handshake deadline and copy bytes without more +application framing. Backpressure comes from the underlying stream. Orderly +EOF is propagated as a half-close; a hard error aborts both directions. + +Legacy QUIC maps each flow to a separate stream in one connection. The TCP +fallback creates one TLS 1.3 connection per flow and has no custom TCP +multiplexer. Protocol v1 carries TCP only. -## Relay phase and shutdown +## Wire migration -After an OK response, both sides clear the protocol-handshake deadline and -copy bytes without further application framing. Backpressure comes from the -underlying stream. An orderly EOF is propagated as a half-close in the other -direction, while a hard error aborts both directions. +Before the Hysteria integration, `--transport=quic` meant AutoCAR v1. It now +means Hysteria v2. Upgrade both UDP endpoints together, or temporarily pin both +sides to the legacy names: -On QUIC, many independent TCP flows share one authenticated QUIC connection, -with one bidirectional QUIC stream per flow. On the TCP fallback, each flow -gets a separate TLS 1.3 connection; protocol v1 does not implement a custom -multiplexer over TCP. +```sh +# Relay +autocar server [common options] --quic-engine=legacy -## Versioning +# Client +autocar client [common options] --transport=legacy-quic +``` -Any incompatible change requires a new version and a new ALPN value. A v1 -implementation must not silently reinterpret unknown kinds, networks, -statuses or non-zero reserved fields. +Only one UDP engine can bind a given address. A staged migration can run the +new Hysteria engine on a second UDP port, then switch clients and finally +retire the old port. The TCP/TLS fallback remains protocol v1, so `auto` can +retain TCP reachability while the UDP endpoints are temporarily mismatched. -This protocol carries TCP only. SOCKS5 UDP ASSOCIATE and QUIC DATAGRAM are not -implemented in v1. +Any incompatible change to AutoCAR v1 requires a new version and ALPN. Hysteria +wire evolution follows its upstream protocol and the pinned module version. diff --git a/docs/assets/autocar-logo.svg b/docs/assets/autocar-logo.svg new file mode 100644 index 0000000..4db3410 --- /dev/null +++ b/docs/assets/autocar-logo.svg @@ -0,0 +1,99 @@ + + + AutoCAR secure dual-ended accelerator + A deep-blue rounded emblem with two converging cyan, blue, and violet acceleration lanes forming a slanted letter A and an arrow, paired with the AutoCAR wordmark. + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + SECURE DUAL-ENDED ACCELERATOR + diff --git a/go.mod b/go.mod index a91f236..5597a7e 100644 --- a/go.mod +++ b/go.mod @@ -2,10 +2,34 @@ module github.com/cppla/autocar go 1.25.0 -require github.com/quic-go/quic-go v0.61.0 +// Local security-hardening fork of Hysteria core v2.12.1. See +// third_party/hysteria-core/AUTOCAR_PATCHES.md. +replace github.com/apernet/hysteria/core/v2 => ./third_party/hysteria-core + +// Local HTTP/3 pre-read admission hook used by the hardened Hysteria core. See +// third_party/quic-go/AUTOCAR_PATCHES.md. +replace github.com/apernet/quic-go => ./third_party/quic-go + +require ( + github.com/apernet/hysteria/core/v2 v2.12.1 + github.com/apernet/hysteria/extras/v2 v2.12.1 + github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e + github.com/quic-go/quic-go v0.61.0 +) require ( + github.com/andybalholm/brotli v1.1.0 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/klauspost/compress v1.18.7 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/quic-go/qpack v0.6.0 // indirect + github.com/refraction-networking/utls v1.8.2 // indirect + github.com/stretchr/objx v0.5.2 // indirect + github.com/stretchr/testify v1.11.1 // indirect golang.org/x/crypto v0.54.0 // indirect - golang.org/x/net v0.56.0 // indirect + golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 // indirect + golang.org/x/net v0.57.0 // indirect golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.40.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index ca322e4..0ac27c4 100644 --- a/go.sum +++ b/go.sum @@ -1,20 +1,47 @@ +github.com/andybalholm/brotli v1.1.0 h1:eLKJA0d02Lf0mVpIDgYnqXcUn0GqVmEFny3VuID1U3M= +github.com/andybalholm/brotli v1.1.0/go.mod h1:sms7XGricyQI9K10gOSf56VKKWS4oLer58Q+mhRPtnY= +github.com/apernet/hysteria/extras/v2 v2.12.1 h1:pLtKedlKSUHGCuUxeVaOOI2UUPlj/5DaqXm01FGrF7U= +github.com/apernet/hysteria/extras/v2 v2.12.1/go.mod h1:QzIFayY1vN8qxX/VpZF+7CGguFpbKDvxbSRp9FmX9V4= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= +github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0= github.com/quic-go/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk= +github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= +github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= github.com/quic-go/quic-go v0.61.0 h1:ui88A53s8MSVYLC56en0KQ17HARk+9986Dn0SBfKNvA= github.com/quic-go/quic-go v0.61.0/go.mod h1:9So2anK4Tp22URSQq00k+Vo2PNkle96ycDPDHL4s9vs= +github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEvV+S9iJ2IdQo= +github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= +github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= +github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= +github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= +github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko= go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o= golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= -golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= -golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= +golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM= +golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= +golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/hy2/auto.go b/internal/hy2/auto.go new file mode 100644 index 0000000..69efb1f --- /dev/null +++ b/internal/hy2/auto.go @@ -0,0 +1,281 @@ +package hy2 + +import ( + "context" + "errors" + "fmt" + "net" + "sync" + "sync/atomic" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +const defaultFallbackCooldown = 30 * time.Second + +const ( + autoRoutePrimary uint32 = iota + autoRouteFallback +) + +type closeDialer interface { + transport.Dialer + Close() error +} + +// AutoConfig configures UDP-first operation with a real TCP/TLS fallback. +type AutoConfig struct { + Primary *Client + Fallback closeDialer + AttemptTimeout time.Duration + Cooldown time.Duration + // OnFallback is called once when a healthy/unknown primary circuit first + // becomes unavailable. It is intended for concise operational logging and + // must not retain secrets or block. + OnFallback func(error) +} + +// AutoClient uses Hysteria v2 first and temporarily routes new TCP flows over +// TLS when the UDP path is unavailable. A valid relay-side destination error +// and an authentication rejection never open the circuit. +type AutoClient struct { + primary *Client + fallback closeDialer + timeout time.Duration + cooldown time.Duration + observer func(error) + route atomic.Uint32 + + fallbackMu sync.RWMutex + lifecycleCtx context.Context + lifecycleCancel context.CancelFunc + + mu sync.Mutex + failedAt time.Time + probing bool + closed bool + closeErr error + closeDone chan struct{} +} + +// NewAutoClient creates an automatic dual-transport dialer. +func NewAutoClient(config AutoConfig) (*AutoClient, error) { + if config.Primary == nil || config.Fallback == nil { + return nil, errors.New("hy2: auto mode requires primary and fallback dialers") + } + if config.AttemptTimeout <= 0 { + return nil, errors.New("hy2: auto attempt timeout must be positive") + } + if config.Cooldown < 0 { + return nil, errors.New("hy2: fallback cooldown cannot be negative") + } + if config.Cooldown == 0 { + config.Cooldown = defaultFallbackCooldown + } + lifecycleCtx, lifecycleCancel := context.WithCancel(context.Background()) + return &AutoClient{ + primary: config.Primary, + fallback: config.Fallback, + timeout: config.AttemptTimeout, + cooldown: config.Cooldown, + observer: config.OnFallback, + lifecycleCtx: lifecycleCtx, + lifecycleCancel: lifecycleCancel, + closeDone: make(chan struct{}), + }, nil +} + +func (c *AutoClient) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + if ctx == nil { + return nil, errors.New("hy2: nil dial context") + } + if c.isClosed() { + return nil, net.ErrClosed + } + if !c.shouldTryPrimary(time.Now()) { + // Close may race between the first closed check and the routing + // decision. Never invoke an already-closed fallback in that window. + if c.isClosed() { + return nil, net.ErrClosed + } + return c.dialFallback(ctx, network, address) + } + primaryCtx, cancel := context.WithTimeout(ctx, c.timeout) + primaryConn, primaryErr := c.primary.DialContext(primaryCtx, network, address) + cancel() + if primaryErr == nil { + c.primarySucceeded() + return primaryConn, nil + } + if callerErr := context.Cause(ctx); callerErr != nil { + c.probeFinished() + return nil, callerErr + } + if IsRemoteDialError(primaryErr) || IsAuthenticationError(primaryErr) { + c.primarySucceeded() + return nil, primaryErr + } + if c.isClosed() { + return nil, net.ErrClosed + } + if c.primaryFailed(time.Now()) && c.observer != nil { + c.observer(primaryErr) + } + fallbackConn, fallbackErr := c.dialFallback(ctx, network, address) + if fallbackErr == nil { + return fallbackConn, nil + } + return nil, errors.Join( + fmt.Errorf("hy2 primary: %w", primaryErr), + fmt.Errorf("TLS fallback: %w", fallbackErr), + ) +} + +func (c *AutoClient) dialFallback(ctx context.Context, network, address string) (net.Conn, error) { + // The read lock linearizes fallback dispatch with Close: once Close marks + // the client closed, no new call can enter the fallback, and an already + // running call receives lifecycle cancellation before Close waits for it. + c.fallbackMu.RLock() + defer c.fallbackMu.RUnlock() + if c.isClosed() { + return nil, net.ErrClosed + } + callCtx, cancel := context.WithCancelCause(ctx) + var stop func() bool + if c.lifecycleCtx != nil { + stop = context.AfterFunc(c.lifecycleCtx, func() { cancel(net.ErrClosed) }) + } + defer func() { + if stop != nil { + stop() + } + cancel(nil) + }() + conn, err := c.fallback.DialContext(callCtx, network, address) + if err == nil { + c.route.Store(autoRouteFallback) + } + return conn, err +} + +// DialPacket opens an accelerated UDP session. TCP/TLS cannot carry SOCKS5 +// UDP without head-of-line blocking, so datagrams intentionally have no +// fallback and report a primary-path failure directly. +func (c *AutoClient) DialPacket(ctx context.Context) (transport.PacketConn, error) { + if ctx == nil { + return nil, errors.New("hy2: nil packet context") + } + if c.isClosed() { + return nil, net.ErrClosed + } + packetCtx, cancel := context.WithTimeout(ctx, c.timeout) + defer cancel() + conn, err := c.primary.DialPacket(packetCtx) + if err == nil { + c.route.Store(autoRoutePrimary) + } + return conn, err +} + +// AccelerationMode reports the transport used by the most recent successful +// flow. A TLS fallback has no QUIC congestion controller. +func (c *AutoClient) AccelerationMode() string { + if c.route.Load() == autoRouteFallback { + return "tls-fallback" + } + return c.primary.AccelerationMode() +} + +// NegotiatedTx reports the primary QUIC path's most recently negotiated +// client-to-server Brutal rate. It is zero while the TLS fallback, BBR or Reno +// is active. +func (c *AutoClient) NegotiatedTx() uint64 { + if c.route.Load() == autoRouteFallback { + return 0 + } + return c.primary.NegotiatedTx() +} + +func (c *AutoClient) isClosed() bool { + c.mu.Lock() + defer c.mu.Unlock() + return c.closed +} + +func (c *AutoClient) shouldTryPrimary(now time.Time) bool { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return false + } + if c.failedAt.IsZero() { + return true + } + if now.Before(c.failedAt.Add(c.cooldown)) || c.probing { + return false + } + c.probing = true + return true +} + +func (c *AutoClient) primarySucceeded() { + c.mu.Lock() + c.failedAt = time.Time{} + c.probing = false + c.mu.Unlock() + c.route.Store(autoRoutePrimary) +} + +func (c *AutoClient) primaryFailed(now time.Time) bool { + c.mu.Lock() + firstFailure := c.failedAt.IsZero() + c.failedAt = now + c.probing = false + c.mu.Unlock() + return firstFailure +} + +func (c *AutoClient) probeFinished() { + c.mu.Lock() + c.probing = false + c.mu.Unlock() +} + +// Close closes both transports. +func (c *AutoClient) Close() error { + c.mu.Lock() + if c.closed { + done := c.closeDone + c.mu.Unlock() + if done != nil { + <-done + } + c.mu.Lock() + err := c.closeErr + c.mu.Unlock() + return err + } + c.closed = true + if c.closeDone == nil { + c.closeDone = make(chan struct{}) + } + done := c.closeDone + lifecycleCancel := c.lifecycleCancel + c.mu.Unlock() + + if lifecycleCancel != nil { + lifecycleCancel() + } + c.fallbackMu.Lock() + err := errors.Join(c.primary.Close(), c.fallback.Close()) + c.fallbackMu.Unlock() + c.mu.Lock() + c.closeErr = err + close(done) + c.mu.Unlock() + return err +} + +var _ transport.Dialer = (*AutoClient)(nil) +var _ transport.PacketDialer = (*AutoClient)(nil) diff --git a/internal/hy2/client.go b/internal/hy2/client.go new file mode 100644 index 0000000..a3ec493 --- /dev/null +++ b/internal/hy2/client.go @@ -0,0 +1,543 @@ +// Package hy2 adapts the Hysteria v2 transport core to AutoCAR's proxy +// interfaces. It provides a long-lived HTTP/3-over-QUIC session, BBR or +// negotiated Brutal congestion control, QUIC datagrams and optional +// Salamander packet obfuscation. +package hy2 + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "strings" + "sync" + "sync/atomic" + "time" + + hyclient "github.com/apernet/hysteria/core/v2/client" + hyerrors "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/extras/v2/obfs" + "github.com/cppla/autocar/internal/protocol" + "github.com/cppla/autocar/internal/transport" +) + +const ( + CongestionBBR = "bbr" + CongestionReno = "reno" + + BBRConservative = "conservative" + BBRStandard = "standard" + BBRAggressive = "aggressive" + + minimumBandwidth = 65536 + maximumBandwidth = 1_000_000_000_000 // 8 Tbit/s; keeps signed QUIC arithmetic safely bounded. + minimumObfsKey = 16 + defaultMaxPendingOpens = 256 +) + +// ClientConfig configures the accelerated Hysteria v2 transport. Bandwidth +// values are bytes per second. Leaving both values at zero selects BBR; +// setting a value selects negotiated Brutal for that sending direction. +type ClientConfig struct { + ServerAddress string + Token string + TLSConfig *tls.Config + + Congestion string + BBRProfile string + MaxTx uint64 + MaxRx uint64 + + DisableLossCompensation bool + FastOpen bool + ObfuscationKey []byte + DisablePathMTUDiscovery bool + DisableGSO bool + DisableChromeParrot bool + MaxIdleTimeout time.Duration + KeepAlivePeriod time.Duration + MaxPendingOpens int +} + +// Client is a reconnecting Hysteria v2 client. The first request establishes +// the authenticated HTTP/3 session lazily. +type Client struct { + config ClientConfig + + coreMu sync.Mutex + core hyclient.Client + attempt *connectAttempt + connectFunc func() (hyclient.Client, *hyclient.HandshakeInfo, error) + + closed sync.Once + closeErr error + closedState atomic.Bool + closeCh chan struct{} + openSlots chan struct{} + + connections atomic.Uint64 + negotiated atomic.Uint64 + udpEnabled atomic.Bool +} + +type connectAttempt struct { + done chan struct{} + core hyclient.Client + err error +} + +// NewClient validates config and creates a lazy reconnecting client. +func NewClient(config ClientConfig) (*Client, error) { + if err := validateClientConfig(config); err != nil { + return nil, err + } + config.Congestion = normalizeCongestion(config.Congestion) + config.BBRProfile = normalizeBBRProfile(config.BBRProfile) + config.ObfuscationKey = append([]byte(nil), config.ObfuscationKey...) + config.TLSConfig = config.TLSConfig.Clone() + if config.MaxPendingOpens == 0 { + config.MaxPendingOpens = defaultMaxPendingOpens + } + + return &Client{ + config: config, + closeCh: make(chan struct{}), + openSlots: make(chan struct{}, config.MaxPendingOpens), + }, nil +} + +func validateClientConfig(config ClientConfig) error { + if config.ServerAddress == "" { + return errors.New("hy2: server address is required") + } + if _, _, err := net.SplitHostPort(config.ServerAddress); err != nil { + return fmt.Errorf("hy2: invalid server address %q: %w", config.ServerAddress, err) + } + if len(config.Token) < protocol.MinTokenLength || len(config.Token) > protocol.MaxTokenLength { + return fmt.Errorf("hy2: token length must be between %d and %d bytes", protocol.MinTokenLength, protocol.MaxTokenLength) + } + if config.TLSConfig == nil { + return errors.New("hy2: TLS config is required") + } + if config.TLSConfig.InsecureSkipVerify { + return errors.New("hy2: InsecureSkipVerify is prohibited") + } + if config.TLSConfig.ServerName == "" || config.TLSConfig.RootCAs == nil { + return errors.New("hy2: verified server name and explicit root CAs are required") + } + if err := validateClientTLSPolicy(config.TLSConfig); err != nil { + return err + } + congestion := normalizeCongestion(config.Congestion) + if congestion != CongestionBBR && congestion != CongestionReno { + return fmt.Errorf("hy2: unsupported congestion controller %q", config.Congestion) + } + profile := normalizeBBRProfile(config.BBRProfile) + if congestion == CongestionBBR && profile != BBRConservative && profile != BBRStandard && profile != BBRAggressive { + return fmt.Errorf("hy2: unsupported BBR profile %q", config.BBRProfile) + } + for name, bandwidth := range map[string]uint64{"MaxTx": config.MaxTx, "MaxRx": config.MaxRx} { + if bandwidth != 0 && bandwidth < minimumBandwidth { + return fmt.Errorf("hy2: %s must be zero or at least %d bytes/s", name, minimumBandwidth) + } + if bandwidth > maximumBandwidth { + return fmt.Errorf("hy2: %s must not exceed %d bytes/s", name, maximumBandwidth) + } + } + if len(config.ObfuscationKey) != 0 && len(config.ObfuscationKey) < minimumObfsKey { + return fmt.Errorf("hy2: obfuscation key must be at least %d bytes", minimumObfsKey) + } + if config.MaxIdleTimeout != 0 && (config.MaxIdleTimeout < 4*time.Second || config.MaxIdleTimeout > 120*time.Second) { + return errors.New("hy2: maximum idle timeout must be zero or between 4s and 120s") + } + if config.KeepAlivePeriod != 0 && (config.KeepAlivePeriod < 2*time.Second || config.KeepAlivePeriod > 60*time.Second) { + return errors.New("hy2: keepalive period must be zero or between 2s and 60s") + } + if config.MaxPendingOpens < 0 || config.MaxPendingOpens > 65536 { + return errors.New("hy2: maximum pending opens must be zero or at most 65536") + } + return nil +} + +func normalizeCongestion(value string) string { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + return CongestionBBR + } + return value +} + +func normalizeBBRProfile(value string) string { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + return BBRStandard + } + return value +} + +func (c *Client) newCoreConfig() (*hyclient.Config, error) { + serverAddress, err := net.ResolveUDPAddr("udp", c.config.ServerAddress) + if err != nil { + return nil, fmt.Errorf("hy2: resolve relay: %w", err) + } + tlsConfig := c.config.TLSConfig.Clone() + getClientCertificate := tlsConfig.GetClientCertificate + if getClientCertificate == nil && len(tlsConfig.Certificates) != 0 { + certificates := append([]tls.Certificate(nil), tlsConfig.Certificates...) + getClientCertificate = func(request *tls.CertificateRequestInfo) (*tls.Certificate, error) { + for i := range certificates { + if err := request.SupportsCertificate(&certificates[i]); err == nil { + return &certificates[i], nil + } + } + return &tls.Certificate{}, nil + } + } + return &hyclient.Config{ + ConnFactory: &packetConnFactory{obfuscationKey: c.config.ObfuscationKey}, + ServerAddr: serverAddress, + Auth: c.config.Token, + TLSConfig: hyclient.TLSConfig{ + ServerName: tlsConfig.ServerName, + InsecureSkipVerify: false, + VerifyPeerCertificate: tlsConfig.VerifyPeerCertificate, + RootCAs: tlsConfig.RootCAs.Clone(), + GetClientCertificate: getClientCertificate, + ECHConfigList: append([]byte(nil), tlsConfig.EncryptedClientHelloConfigList...), + }, + QUICConfig: hyclient.QUICConfig{ + MaxIdleTimeout: c.config.MaxIdleTimeout, + KeepAlivePeriod: c.config.KeepAlivePeriod, + DisablePathMTUDiscovery: c.config.DisablePathMTUDiscovery, + DisableGSO: c.config.DisableGSO, + DisableChromeParrot: c.config.DisableChromeParrot, + }, + CongestionConfig: hyclient.CongestionConfig{ + Type: c.config.Congestion, + BBRProfile: c.config.BBRProfile, + }, + BandwidthConfig: hyclient.BandwidthConfig{ + MaxTx: c.config.MaxTx, + MaxRx: c.config.MaxRx, + DisableLossCompensation: c.config.DisableLossCompensation, + }, + FastOpen: c.config.FastOpen, + }, nil +} + +// coreForContext returns the active authenticated session. Connection setup is +// a context-aware single flight: a UDP black hole can leave at most one +// bounded upstream handshake running, while all other callers wait without +// spawning their own reconnect attempts. A successful session is shared by +// all streams and datagram associations. +func (c *Client) coreForContext(ctx context.Context) (hyclient.Client, error) { + c.coreMu.Lock() + if c.closedState.Load() { + c.coreMu.Unlock() + return nil, net.ErrClosed + } + if c.core != nil { + core := c.core + c.coreMu.Unlock() + return core, nil + } + attempt := c.attempt + if attempt == nil { + attempt = &connectAttempt{done: make(chan struct{})} + c.attempt = attempt + go c.connect(attempt) + } + c.coreMu.Unlock() + + select { + case <-attempt.done: + if attempt.err != nil { + return nil, attempt.err + } + return attempt.core, nil + case <-ctx.Done(): + return nil, context.Cause(ctx) + case <-c.closeCh: + return nil, net.ErrClosed + } +} + +func (c *Client) connect(attempt *connectAttempt) { + connect := c.connectFunc + if connect == nil { + connect = func() (hyclient.Client, *hyclient.HandshakeInfo, error) { + config, err := c.newCoreConfig() + if err != nil { + return nil, nil, err + } + return hyclient.NewClient(config) + } + } + core, info, err := connect() + if err == nil && (core == nil || info == nil) { + err = errors.New("hy2: connector returned an incomplete session") + } + if err != nil { + err = c.addChromeCertificateHint(err) + err = fmt.Errorf("hy2: establish authenticated session: %w", err) + } + + c.coreMu.Lock() + if err == nil && c.closedState.Load() { + err = net.ErrClosed + } + if err == nil { + c.core = core + c.connections.Add(1) + c.negotiated.Store(info.Tx) + c.udpEnabled.Store(info.UDPEnabled) + } + attempt.core = core + attempt.err = err + if c.attempt == attempt { + c.attempt = nil + } + close(attempt.done) + c.coreMu.Unlock() + + if err != nil && core != nil { + _ = core.Close() + } +} + +// addChromeCertificateHint preserves the handshake error while explaining a +// known compatibility constraint of the Chrome-parroting ClientHello. Chrome's +// advertised signature schemes intentionally omit Ed25519; an operator can +// either serve an ECDSA P-256/P-384/RSA certificate or explicitly disable parroting. +func (c *Client) addChromeCertificateHint(err error) error { + if err == nil || c.config.DisableChromeParrot { + return err + } + message := strings.ToLower(err.Error()) + if !strings.Contains(message, "handshake failure") && + !strings.Contains(message, "signature algorithm") { + return err + } + return fmt.Errorf("%w (Chrome QUIC fingerprinting is enabled; if the relay certificate is Ed25519, use ECDSA P-256/P-384/RSA or pass --disable-chrome-parrot)", err) +} + +func (c *Client) invalidate(core hyclient.Client, err error) { + var closedError hyerrors.ClosedError + if !errors.As(err, &closedError) { + return + } + c.coreMu.Lock() + if c.core != core { + c.coreMu.Unlock() + return + } + c.core = nil + c.udpEnabled.Store(false) + c.negotiated.Store(0) + c.coreMu.Unlock() + _ = core.Close() +} + +type packetConnFactory struct { + obfuscationKey []byte +} + +func (f *packetConnFactory) New(net.Addr) (net.PacketConn, error) { + conn, err := net.ListenUDP("udp", nil) + if err != nil { + return nil, err + } + if len(f.obfuscationKey) == 0 { + return conn, nil + } + wrapped, err := obfs.WrapPacketConnSalamander(conn, f.obfuscationKey) + if err != nil { + _ = conn.Close() + return nil, err + } + return wrapped, nil +} + +type tcpResult struct { + conn net.Conn + err error +} + +// DialContext opens a TCP stream over the authenticated QUIC session. +func (c *Client) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + if ctx == nil { + return nil, errors.New("hy2: nil dial context") + } + if c.closedState.Load() { + return nil, net.ErrClosed + } + if network != "tcp" { + return nil, fmt.Errorf("hy2: unsupported network %q", network) + } + if _, _, err := net.SplitHostPort(address); err != nil { + return nil, fmt.Errorf("hy2: invalid destination %q: %w", address, err) + } + if err := c.acquireOpen(ctx); err != nil { + return nil, err + } + core, err := c.coreForContext(ctx) + if err != nil { + c.releaseOpen() + return nil, err + } + result := make(chan tcpResult) + go func() { + defer c.releaseOpen() + conn, err := core.TCP(address) + c.invalidate(core, err) + value := tcpResult{conn: conn, err: err} + select { + case result <- value: + case <-ctx.Done(): + if conn != nil { + _ = conn.Close() + } + case <-c.closeCh: + if conn != nil { + _ = conn.Close() + } + } + }() + select { + case value := <-result: + if value.err != nil { + return nil, fmt.Errorf("hy2: open TCP stream: %w", value.err) + } + if err := context.Cause(ctx); err != nil { + _ = value.conn.Close() + return nil, err + } + return value.conn, nil + case <-ctx.Done(): + return nil, context.Cause(ctx) + case <-c.closeCh: + return nil, net.ErrClosed + } +} + +func (c *Client) acquireOpen(ctx context.Context) error { + if c.openSlots == nil { + return nil + } + select { + case c.openSlots <- struct{}{}: + return nil + case <-ctx.Done(): + return context.Cause(ctx) + case <-c.closeCh: + return net.ErrClosed + } +} + +func (c *Client) releaseOpen() { + if c.openSlots != nil { + <-c.openSlots + } +} + +// DialPacket opens one logical UDP session carried by QUIC DATAGRAM frames. +func (c *Client) DialPacket(ctx context.Context) (transport.PacketConn, error) { + if ctx == nil { + return nil, errors.New("hy2: nil packet context") + } + if c.closedState.Load() { + return nil, net.ErrClosed + } + core, err := c.coreForContext(ctx) + if err != nil { + return nil, err + } + conn, err := core.UDP() + c.invalidate(core, err) + if err != nil { + return nil, fmt.Errorf("hy2: open UDP session: %w", err) + } + if err := context.Cause(ctx); err != nil { + _ = conn.Close() + return nil, err + } + return &packetConn{core: conn}, nil +} + +type packetConn struct { + core hyclient.HyUDPConn +} + +func (c *packetConn) Send(payload []byte, address string) error { + return c.core.Send(payload, address) +} + +func (c *packetConn) Receive() ([]byte, string, error) { + return c.core.Receive() +} + +func (c *packetConn) Close() error { return c.core.Close() } + +func (c *packetConn) MaxPayloadSize() int { return hyclient.MaxUDPSize } + +// NegotiatedTx reports the most recently negotiated client-to-server Brutal +// rate in bytes/s. Zero means BBR or Reno is active for that direction. +func (c *Client) NegotiatedTx() uint64 { return c.negotiated.Load() } + +// UDPEnabled reports whether the current server handshake enabled datagrams. +func (c *Client) UDPEnabled() bool { return c.udpEnabled.Load() } + +// ConnectionCount reports successful authenticated session establishments, +// including reconnects. +func (c *Client) ConnectionCount() uint64 { return c.connections.Load() } + +// AccelerationMode reports the active outbound congestion-control mode. A +// non-zero negotiated rate always means Brutal; otherwise the configured BBR +// profile or Reno controls the connection. +func (c *Client) AccelerationMode() string { + if c.negotiated.Load() > 0 { + return "brutal" + } + if normalizeCongestion(c.config.Congestion) == CongestionReno { + return CongestionReno + } + return CongestionBBR + "-" + normalizeBBRProfile(c.config.BBRProfile) +} + +// Close permanently closes the reconnecting client. +func (c *Client) Close() error { + c.closed.Do(func() { + c.closedState.Store(true) + if c.closeCh != nil { + close(c.closeCh) + } + c.coreMu.Lock() + core := c.core + c.core = nil + c.coreMu.Unlock() + if core != nil { + c.closeErr = core.Close() + } + }) + return c.closeErr +} + +// IsRemoteDialError reports errors returned by an authenticated relay after it +// attempted the requested target. Auto mode must not retry those requests via +// another transport, because the primary path itself is healthy. +func IsRemoteDialError(err error) bool { + var dialError hyerrors.DialError + return errors.As(err, &dialError) +} + +// IsAuthenticationError reports Hysteria authentication rejection. +func IsAuthenticationError(err error) bool { + var authError hyerrors.AuthError + return errors.As(err, &authError) +} + +var _ transport.Dialer = (*Client)(nil) +var _ transport.PacketDialer = (*Client)(nil) +var _ transport.PacketConn = (*packetConn)(nil) +var _ transport.PacketPayloadSizer = (*packetConn)(nil) diff --git a/internal/hy2/hy2_test.go b/internal/hy2/hy2_test.go new file mode 100644 index 0000000..6d535ff --- /dev/null +++ b/internal/hy2/hy2_test.go @@ -0,0 +1,1614 @@ +package hy2 + +import ( + "context" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "errors" + "io" + "net" + "net/http" + "net/netip" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + hyclient "github.com/apernet/hysteria/core/v2/client" + hyserver "github.com/apernet/hysteria/core/v2/server" + "github.com/apernet/quic-go/http3" + "github.com/cppla/autocar/internal/security" + "github.com/cppla/autocar/internal/transport" +) + +const ( + testToken = "correct horse battery staple" + wrongTestToken = "incorrect horse battery staple" +) + +func TestBBRLoopbackTCPEcho(t *testing.T) { + target := startTCPEcho(t) + server, clientTLS, outbound := startTestServer(t, nil) + client := newTestClient(t, server, clientTLS, nil) + + exchangeTCP(t, client, target, "BBR keeps one authenticated QUIC session hot") + if got := client.NegotiatedTx(); got != 0 { + t.Fatalf("NegotiatedTx = %d, want 0 while BBR is active", got) + } + if got := server.admission.lastTx.Load(); got != 0 { + t.Fatalf("server negotiated Tx = %d, want 0 while BBR is active", got) + } + if got := client.ConnectionCount(); got != 1 { + t.Fatalf("authenticated connection count = %d, want 1", got) + } + if got := client.AccelerationMode(); got != "bbr-standard" { + t.Fatalf("acceleration mode = %q, want bbr-standard", got) + } + if got := outbound.tcpDials.Load(); got != 1 { + t.Fatalf("target dial count = %d, want 1", got) + } +} + +func TestBrutalNegotiatesBothDirectionsAndAppliesServerCaps(t *testing.T) { + const ( + serverMaxTx = 300_000 + serverMaxRx = 400_000 + clientMaxTx = 700_000 + clientMaxRx = 600_000 + ) + target := startTCPEcho(t) + server, clientTLS, _ := startTestServer(t, func(config *ServerConfig) { + config.MaxTx = serverMaxTx + config.MaxRx = serverMaxRx + config.AllowClientBandwidth = true + }) + client := newTestClient(t, server, clientTLS, func(config *ClientConfig) { + config.MaxTx = clientMaxTx + config.MaxRx = clientMaxRx + }) + + exchangeTCP(t, client, target, "Brutal is negotiated independently in both directions") + if got := client.NegotiatedTx(); got != serverMaxRx { + t.Fatalf("client Tx = %d, want server Rx cap %d", got, serverMaxRx) + } + if got := server.admission.lastTx.Load(); got != serverMaxTx { + t.Fatalf("server Tx = %d, want server Tx cap %d", got, serverMaxTx) + } + if got := client.AccelerationMode(); got != "brutal" { + t.Fatalf("acceleration mode = %q, want brutal", got) + } +} + +func TestServerIgnoresClientBandwidthByDefault(t *testing.T) { + const requested = 900_000 + target := startTCPEcho(t) + server, clientTLS, _ := startTestServer(t, nil) + client := newTestClient(t, server, clientTLS, func(config *ClientConfig) { + config.MaxTx = requested + config.MaxRx = requested + }) + + exchangeTCP(t, client, target, "secure default keeps the model controller") + if got := client.NegotiatedTx(); got != 0 { + t.Fatalf("client Tx = %d, want BBR negotiation value 0", got) + } + if got := server.admission.lastTx.Load(); got != 0 { + t.Fatalf("server Tx = %d, want BBR negotiation value 0", got) + } + if got := client.AccelerationMode(); got != "bbr-standard" { + t.Fatalf("acceleration mode = %q, want bbr-standard", got) + } +} + +func TestAccelerationModeLabelsControllerAndProfile(t *testing.T) { + for _, test := range []struct { + name string + config ClientConfig + want string + }{ + {name: "default BBR", want: "bbr-standard"}, + {name: "conservative BBR", config: ClientConfig{Congestion: CongestionBBR, BBRProfile: BBRConservative}, want: "bbr-conservative"}, + {name: "aggressive BBR", config: ClientConfig{Congestion: CongestionBBR, BBRProfile: BBRAggressive}, want: "bbr-aggressive"}, + {name: "Reno", config: ClientConfig{Congestion: CongestionReno}, want: "reno"}, + } { + t.Run(test.name, func(t *testing.T) { + client := &Client{config: test.config} + if got := client.AccelerationMode(); got != test.want { + t.Fatalf("AccelerationMode = %q, want %q", got, test.want) + } + }) + } + client := &Client{config: ClientConfig{Congestion: CongestionBBR, BBRProfile: BBRStandard}} + auto := &AutoClient{primary: client} + if got := auto.AccelerationMode(); got != "bbr-standard" { + t.Fatalf("AutoClient AccelerationMode = %q", got) + } + client.negotiated.Store(123_456) + if got := auto.NegotiatedTx(); got != 123_456 { + t.Fatalf("AutoClient NegotiatedTx = %d, want 123456", got) + } + auto.route.Store(autoRouteFallback) + if got := auto.AccelerationMode(); got != "tls-fallback" { + t.Fatalf("fallback AccelerationMode = %q", got) + } + if got := auto.NegotiatedTx(); got != 0 { + t.Fatalf("fallback NegotiatedTx = %d, want 0", got) + } + auto.primarySucceeded() + if got := auto.AccelerationMode(); got != "brutal" { + t.Fatalf("restored primary AccelerationMode = %q", got) + } +} + +func TestChromeHandshakeFailureIncludesCertificateGuidance(t *testing.T) { + client := &Client{ + config: ClientConfig{}, + closeCh: make(chan struct{}), + connectFunc: func() (hyclient.Client, *hyclient.HandshakeInfo, error) { + return nil, nil, errors.New("remote error: tls: handshake failure") + }, + } + _, err := client.coreForContext(context.Background()) + if err == nil || !strings.Contains(err.Error(), "Ed25519") || !strings.Contains(err.Error(), "--disable-chrome-parrot") { + t.Fatalf("Chrome handshake guidance error = %v", err) + } + + disabled := &Client{ + config: ClientConfig{DisableChromeParrot: true}, + closeCh: make(chan struct{}), + connectFunc: func() (hyclient.Client, *hyclient.HandshakeInfo, error) { + return nil, nil, errors.New("remote error: tls: handshake failure") + }, + } + _, err = disabled.coreForContext(context.Background()) + if err == nil || strings.Contains(err.Error(), "Ed25519") { + t.Fatalf("disabled Chrome mode added misleading guidance: %v", err) + } +} + +func TestWrongTokenNeverDialsTarget(t *testing.T) { + server, clientTLS, outbound := startTestServer(t, nil) + client := newTestClient(t, server, clientTLS, func(config *ClientConfig) { + config.Token = wrongTestToken + }) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + _, err := client.DialContext(ctx, "tcp", "127.0.0.1:1") + if err == nil || !IsAuthenticationError(err) { + t.Fatalf("wrong-token error = %v, want authentication rejection", err) + } + if got := outbound.tcpDials.Load(); got != 0 { + t.Fatalf("wrong token caused %d target dials", got) + } +} + +func TestWrongCAPreventsSessionAndTargetDial(t *testing.T) { + server, _, outbound := startTestServer(t, nil) + _, wrongClientTLS := testTLSConfigs(t) + client := newTestClient(t, server, wrongClientTLS, nil) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + _, err := client.DialContext(ctx, "tcp", "127.0.0.1:1") + if err == nil { + t.Fatal("relay certificate signed by an untrusted CA was accepted") + } + if got := outbound.tcpDials.Load(); got != 0 { + t.Fatalf("failed TLS verification caused %d target dials", got) + } +} + +func TestClientVerifyPeerCertificateCallbackIsPreserved(t *testing.T) { + server, clientTLS, outbound := startTestServer(t, nil) + var called atomic.Bool + client := newTestClient(t, server, clientTLS, func(config *ClientConfig) { + config.TLSConfig = config.TLSConfig.Clone() + config.TLSConfig.VerifyPeerCertificate = func([][]byte, [][]*x509.Certificate) error { + called.Store(true) + return errors.New("test certificate policy rejected the relay") + } + }) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if _, err := client.DialContext(ctx, "tcp", "127.0.0.1:1"); err == nil { + t.Fatal("custom certificate policy was silently ignored") + } + if !called.Load() { + t.Fatal("VerifyPeerCertificate was not called") + } + if got := outbound.tcpDials.Load(); got != 0 { + t.Fatalf("rejected TLS policy caused %d target dials", got) + } +} + +func TestHysteriaStrictMutualTLS(t *testing.T) { + target := startTCPEcho(t) + var clientCertificate tls.Certificate + server, clientTLS, outbound := startTestServer(t, func(config *ServerConfig) { + clientCertificate = config.TLSConfig.Certificates[0] + leaf, err := x509.ParseCertificate(clientCertificate.Certificate[0]) + if err != nil { + t.Fatal(err) + } + clientCAs := x509.NewCertPool() + clientCAs.AddCert(leaf) + config.TLSConfig.ClientCAs = clientCAs + config.TLSConfig.ClientAuth = tls.RequireAndVerifyClientCert + }) + + withoutCertificate := newTestClient(t, server, clientTLS, nil) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + _, err := withoutCertificate.DialContext(ctx, "tcp", target) + cancel() + if err == nil { + t.Fatal("strict mTLS accepted a client without a certificate") + } + if got := outbound.tcpDials.Load(); got != 0 { + t.Fatalf("client without an mTLS certificate caused %d target dials", got) + } + + withCertificate := newTestClient(t, server, clientTLS, func(config *ClientConfig) { + config.TLSConfig = config.TLSConfig.Clone() + config.TLSConfig.Certificates = []tls.Certificate{clientCertificate} + }) + exchangeTCP(t, withCertificate, target, "strict Hysteria mutual TLS") +} + +func TestClientRejectsTLSVerificationPoliciesTheAdapterCannotPreserve(t *testing.T) { + _, base := testTLSConfigs(t) + tests := []struct { + name string + field string + modify func(*tls.Config) + }{ + { + name: "VerifyConnection", + field: "VerifyConnection", + modify: func(config *tls.Config) { + config.VerifyConnection = func(tls.ConnectionState) error { return nil } + }, + }, + { + name: "ECH rejection verifier", + field: "EncryptedClientHelloRejectionVerify", + modify: func(config *tls.Config) { + config.EncryptedClientHelloRejectionVerify = func(tls.ConnectionState) error { return nil } + }, + }, + {name: "custom time", field: "Time", modify: func(config *tls.Config) { config.Time = time.Now }}, + {name: "custom randomness", field: "Rand", modify: func(config *tls.Config) { config.Rand = rand.Reader }}, + { + name: "curve policy", + field: "CurvePreferences", + modify: func(config *tls.Config) { config.CurvePreferences = []tls.CurveID{tls.CurveP256} }, + }, + { + name: "TLS 1.2 maximum", + field: "MaxVersion", + modify: func(config *tls.Config) { config.MaxVersion = tls.VersionTLS12 }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + config := base.Clone() + test.modify(config) + _, err := NewClient(ClientConfig{ + ServerAddress: "127.0.0.1:443", + Token: testToken, + TLSConfig: config, + }) + if err == nil || !strings.Contains(err.Error(), test.field) { + t.Fatalf("error = %v, want explicit %s rejection", err, test.field) + } + }) + } +} + +func TestServerRejectsTLSPoliciesTheAdapterCannotPreserve(t *testing.T) { + base, _ := testTLSConfigs(t) + tests := []struct { + name string + field string + modify func(*tls.Config) + }{ + { + name: "GetConfigForClient", + field: "GetConfigForClient", + modify: func(config *tls.Config) { + config.GetConfigForClient = func(*tls.ClientHelloInfo) (*tls.Config, error) { return config, nil } + }, + }, + { + name: "VerifyConnection", + field: "VerifyConnection", + modify: func(config *tls.Config) { + config.VerifyConnection = func(tls.ConnectionState) error { return nil } + }, + }, + { + name: "VerifyPeerCertificate", + field: "VerifyPeerCertificate", + modify: func(config *tls.Config) { + config.VerifyPeerCertificate = func([][]byte, [][]*x509.Certificate) error { return nil } + }, + }, + {name: "custom time", field: "Time", modify: func(config *tls.Config) { config.Time = time.Now }}, + {name: "custom randomness", field: "Rand", modify: func(config *tls.Config) { config.Rand = rand.Reader }}, + { + name: "curve policy", + field: "CurvePreferences", + modify: func(config *tls.Config) { config.CurvePreferences = []tls.CurveID{tls.CurveP256} }, + }, + { + name: "certificate name map", + field: "NameToCertificate", + modify: func(config *tls.Config) { + config.NameToCertificate = map[string]*tls.Certificate{"relay.example": &config.Certificates[0]} + }, + }, + { + name: "session tickets disabled", + field: "SessionTicketsDisabled", + modify: func(config *tls.Config) { config.SessionTicketsDisabled = true }, + }, + { + name: "custom session ticket key", + field: "SessionTicketKey", + modify: func(config *tls.Config) { config.SessionTicketKey[0] = 1 }, + }, + { + name: "custom session wrapper", + field: "WrapSession", + modify: func(config *tls.Config) { + config.WrapSession = func(tls.ConnectionState, *tls.SessionState) ([]byte, error) { return nil, nil } + }, + }, + { + name: "custom session unwrapper", + field: "UnwrapSession", + modify: func(config *tls.Config) { + config.UnwrapSession = func([]byte, tls.ConnectionState) (*tls.SessionState, error) { return nil, nil } + }, + }, + { + name: "client auth without CA", + field: "ClientAuth", + modify: func(config *tls.Config) { config.ClientAuth = tls.RequireAnyClientCert }, + }, + { + name: "client CA without strict auth", + field: "ClientCAs", + modify: func(config *tls.Config) { config.ClientCAs = x509.NewCertPool() }, + }, + { + name: "TLS 1.2 maximum", + field: "MaxVersion", + modify: func(config *tls.Config) { config.MaxVersion = tls.VersionTLS12 }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + config := base.Clone() + test.modify(config) + err := validateServerConfig(ServerConfig{ + Address: "127.0.0.1:0", + Token: testToken, + TLSConfig: config, + Outbound: &testOutbound{}, + }) + if err == nil || !strings.Contains(err.Error(), test.field) { + t.Fatalf("error = %v, want explicit %s rejection", err, test.field) + } + }) + } +} + +func TestTLSAdapterAllowsTLS13IrrelevantAndForwardedCLIFields(t *testing.T) { + serverTLS, clientTLS := testTLSConfigs(t) + clientTLS.NextProtos = []string{security.ALPN} + clientTLS.CipherSuites = []uint16{tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256} + clientTLS.DynamicRecordSizingDisabled = true + clientTLS.Renegotiation = tls.RenegotiateFreelyAsClient + clientTLS.ClientSessionCache = tls.NewLRUClientSessionCache(8) + if _, err := NewClient(ClientConfig{ + ServerAddress: "127.0.0.1:443", + Token: testToken, + TLSConfig: clientTLS, + }); err != nil { + t.Fatalf("ordinary TLS 1.3 client config was rejected: %v", err) + } + + serverTLS.NextProtos = []string{security.ALPN} + serverTLS.CipherSuites = []uint16{tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256} + serverTLS.DynamicRecordSizingDisabled = true + serverTLS.Renegotiation = tls.RenegotiateNever + if err := validateServerConfig(ServerConfig{ + Address: "127.0.0.1:0", + Token: testToken, + TLSConfig: serverTLS, + Outbound: &testOutbound{}, + }); err != nil { + t.Fatalf("ordinary TLS 1.3 server config was rejected: %v", err) + } +} + +func TestHTTP3CoverLooksLikeOrdinaryWebsite(t *testing.T) { + server, clientTLS, _ := startTestServer(t, func(config *ServerConfig) { + config.MasqueradeHandler = NewCoverHandler("AutoCAR Edge") + }) + h3 := &http3.Transport{TLSClientConfig: clientTLS.Clone()} + t.Cleanup(func() { _ = h3.Close() }) + httpClient := &http.Client{Transport: h3, Timeout: 3 * time.Second} + + response, err := httpClient.Get("https://" + server.Addr().String() + "/") + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, err := io.ReadAll(response.Body) + if err != nil { + t.Fatal(err) + } + if response.StatusCode != http.StatusOK || !strings.Contains(string(body), "AutoCAR Edge") { + t.Fatalf("cover response status=%d body=%q", response.StatusCode, body) + } + if got := response.Header.Get("X-Content-Type-Options"); got != "nosniff" { + t.Fatalf("X-Content-Type-Options = %q", got) + } +} + +func TestHTTP3CoverRejectsOversizedRequestHeaders(t *testing.T) { + server, clientTLS, _ := startTestServer(t, nil) + h3 := &http3.Transport{TLSClientConfig: clientTLS.Clone()} + t.Cleanup(func() { _ = h3.Close() }) + httpClient := &http.Client{Transport: h3, Timeout: 3 * time.Second} + request, err := http.NewRequest(http.MethodGet, "https://"+server.Addr().String()+"/", nil) + if err != nil { + t.Fatal(err) + } + request.Header.Set("X-Oversized", strings.Repeat("a", defaultMaxHTTPHeaderBytes*2)) + response, err := httpClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusRequestHeaderFieldsTooLarge { + t.Fatalf("oversized request status = %d, want %d", response.StatusCode, http.StatusRequestHeaderFieldsTooLarge) + } +} + +func TestSalamanderMatchingKeyWorksAndWrongKeyFails(t *testing.T) { + key := []byte("salamander integration secret") + target := startTCPEcho(t) + server, clientTLS, _ := startTestServer(t, func(config *ServerConfig) { + config.ObfuscationKey = key + }) + client := newTestClient(t, server, clientTLS, func(config *ClientConfig) { + config.ObfuscationKey = key + }) + exchangeTCP(t, client, target, "matching Salamander PSKs") + + badClient, err := NewClient(ClientConfig{ + ServerAddress: server.Addr().String(), + Token: testToken, + TLSConfig: clientTLS, + ObfuscationKey: []byte("different integration secret"), + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(context.Background(), 250*time.Millisecond) + defer cancel() + if _, err := badClient.DialContext(ctx, "tcp", target); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("wrong Salamander key error = %v, want deadline exceeded", err) + } + // The Hysteria core does not expose a handshake context. DialContext still + // returns promptly; close asynchronously so the core's bounded handshake + // timeout can release its reconnect lock without slowing this test. + go func() { _ = badClient.Close() }() +} + +func TestUDPDatagramBidirectionalLoopback(t *testing.T) { + target := startUDPEcho(t) + server, clientTLS, outbound := startTestServer(t, nil) + client := newTestClient(t, server, clientTLS, nil) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + packet, err := client.DialPacket(ctx) + if err != nil { + t.Fatal(err) + } + defer packet.Close() + if !client.UDPEnabled() { + t.Fatal("server handshake did not enable QUIC DATAGRAM") + } + + payload := []byte("UDP survives without TCP head-of-line blocking") + if err := packet.Send(payload, target); err != nil { + t.Fatal(err) + } + type receiveResult struct { + payload []byte + address string + err error + } + result := make(chan receiveResult, 1) + go func() { + data, address, err := packet.Receive() + result <- receiveResult{payload: data, address: address, err: err} + }() + select { + case value := <-result: + if value.err != nil { + t.Fatal(value.err) + } + if string(value.payload) != string(payload) { + t.Fatalf("UDP payload = %q, want %q", value.payload, payload) + } + if value.address != target { + t.Fatalf("UDP source = %q, want %q", value.address, target) + } + case <-ctx.Done(): + t.Fatalf("UDP receive: %v", context.Cause(ctx)) + } + if outbound.udpChecks.Load() == 0 { + t.Fatal("UDP destination policy was not checked") + } +} + +func TestSafeUDPOutboundChecksInitialAddressAndRebinding(t *testing.T) { + unsafeResolver := &hyRotatingResolver{answers: [][]netip.Addr{{netip.MustParseAddr("169.254.169.254")}}} + unsafeDialer := security.NewSafeDialer(security.SafeDialerOptions{Resolver: unsafeResolver}) + unsafeOutbound := &safeOutbound{dialer: unsafeDialer, timeout: time.Second} + if _, err := unsafeOutbound.UDP("metadata.example:53"); !errors.Is(err, security.ErrUnsafeAddress) { + t.Fatalf("initial unsafe UDP address error = %v", err) + } + + rebindingResolver := &hyRotatingResolver{answers: [][]netip.Addr{ + {netip.MustParseAddr("8.8.8.8")}, + {netip.MustParseAddr("127.0.0.1")}, + }} + rebindingDialer := security.NewSafeDialer(security.SafeDialerOptions{Resolver: rebindingResolver}) + rebindingOutbound := &safeOutbound{dialer: rebindingDialer, timeout: time.Second} + conn, err := rebindingOutbound.UDP("rebinding.example:53") + if err != nil { + t.Fatal(err) + } + defer conn.Close() + if _, err := conn.WriteTo([]byte("must not escape"), "rebinding.example:53"); !errors.Is(err, security.ErrUnsafeAddress) { + t.Fatalf("rebound UDP destination error = %v", err) + } + if got := rebindingResolver.calls.Load(); got != 2 { + t.Fatalf("resolver calls = %d, want initial and send-time checks", got) + } +} + +func TestSafeOutboundGloballyLimitsTCPAndUDPSessions(t *testing.T) { + dialer := security.NewSafeDialer(security.SafeDialerOptions{ + Dialer: transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + local, peer := net.Pipe() + _ = peer.Close() + return local, nil + }), + }) + outbound := &safeOutbound{ + dialer: dialer, timeout: time.Second, + tcpSlots: make(chan struct{}, 1), + udpSlots: make(chan struct{}, 1), + } + + tcp, err := outbound.TCP("8.8.8.8:443") + if err != nil { + t.Fatal(err) + } + if _, err := outbound.TCP("8.8.8.8:443"); !errors.Is(err, ErrOutboundCapacity) { + t.Fatalf("second TCP error = %v, want capacity rejection", err) + } + if err := tcp.Close(); err != nil { + t.Fatal(err) + } + if err := tcp.Close(); err != nil { + t.Fatal(err) + } + tcp, err = outbound.TCP("8.8.8.8:443") + if err != nil { + t.Fatalf("TCP slot was not released: %v", err) + } + _ = tcp.Close() + + udp, err := outbound.UDP("8.8.8.8:53") + if err != nil { + t.Fatal(err) + } + if _, err := outbound.UDP("8.8.8.8:53"); !errors.Is(err, ErrOutboundCapacity) { + t.Fatalf("second UDP error = %v, want capacity rejection", err) + } + if err := udp.Close(); err != nil { + t.Fatal(err) + } + if err := udp.Close(); err != nil { + t.Fatal(err) + } + udp, err = outbound.UDP("8.8.8.8:53") + if err != nil { + t.Fatalf("UDP slot was not released: %v", err) + } + _ = udp.Close() +} + +func TestSafeUDPConnDropsUnsolicitedSources(t *testing.T) { + relay, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer relay.Close() + approved, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer approved.Close() + unsolicited, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer unsolicited.Close() + + conn := &safeUDPConn{conn: relay} + approvedAddress := approved.LocalAddr().(*net.UDPAddr).AddrPort() + approvedAddress = netip.AddrPortFrom(approvedAddress.Addr().Unmap(), approvedAddress.Port()) + if _, err := conn.writeToDestination([]byte("authorize"), approvedAddress); err != nil { + t.Fatalf("authorize destination: %v", err) + } + target := relay.LocalAddr().(*net.UDPAddr).AddrPort() + if _, err := unsolicited.WriteToUDPAddrPort([]byte("injected"), target); err != nil { + t.Fatal(err) + } + if _, err := approved.WriteToUDPAddrPort([]byte("approved"), target); err != nil { + t.Fatal(err) + } + if err := relay.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + buffer := make([]byte, 32) + n, source, err := conn.ReadFrom(buffer) + if err != nil { + t.Fatal(err) + } + if got := string(buffer[:n]); got != "approved" { + t.Fatalf("payload = %q, want approved", got) + } + if source != approvedAddress.String() { + t.Fatalf("source = %q, want %q", source, approvedAddress) + } +} + +func TestSafeUDPConnDestinationCapacity(t *testing.T) { + relay, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer relay.Close() + conn := &safeUDPConn{conn: relay} + host := netip.MustParseAddr("127.0.0.1") + + for i := range maxUDPAllowedDestinations { + destination := netip.AddrPortFrom(host, uint16(20_000+i)) + if _, err := conn.writeToDestination([]byte("fill"), destination); err != nil { + t.Fatalf("authorize destination %d: %v", i, err) + } + } + newDestination := netip.AddrPortFrom(host, 30_000) + if _, err := conn.writeToDestination([]byte("reject"), newDestination); !errors.Is(err, ErrUDPDestinationCapacity) { + t.Fatalf("new destination after capacity error = %v, want %v", err, ErrUDPDestinationCapacity) + } + if conn.destinationAllowed(newDestination) { + t.Fatal("capacity-rejected destination was authorized") + } + + // Filling the set must not revoke destinations that were already approved. + existing := netip.AddrPortFrom(host, 20_000) + if _, err := conn.writeToDestination([]byte("existing"), existing); err != nil { + t.Fatalf("existing destination after capacity: %v", err) + } + conn.allowedMu.RLock() + allowedCount := len(conn.allowed) + conn.allowedMu.RUnlock() + if allowedCount != maxUDPAllowedDestinations { + t.Fatalf("allowed destination count = %d, want %d", allowedCount, maxUDPAllowedDestinations) + } +} + +func TestSafeUDPConnFailedWriteDoesNotAuthorize(t *testing.T) { + relay, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + if err := relay.Close(); err != nil { + t.Fatal(err) + } + conn := &safeUDPConn{conn: relay} + destination := netip.MustParseAddrPort("127.0.0.1:20000") + if _, err := conn.writeToDestination([]byte("must fail"), destination); err == nil { + t.Fatal("write on closed socket unexpectedly succeeded") + } + if conn.destinationAllowed(destination) { + t.Fatal("failed write authorized its destination") + } +} + +func TestSafeUDPConnDestinationSetConcurrent(t *testing.T) { + relay, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer relay.Close() + conn := &safeUDPConn{conn: relay} + host := netip.MustParseAddr("127.0.0.1") + const attempts = maxUDPAllowedDestinations + 64 + + var wait sync.WaitGroup + start := make(chan struct{}) + errorsByAttempt := make(chan error, attempts) + for i := range attempts { + destination := netip.AddrPortFrom(host, uint16(31_000+i)) + wait.Add(1) + go func() { + defer wait.Done() + <-start + _, err := conn.writeToDestination([]byte("concurrent"), destination) + _ = conn.destinationAllowed(destination) + errorsByAttempt <- err + }() + } + close(start) + wait.Wait() + close(errorsByAttempt) + + var successes int + for err := range errorsByAttempt { + switch { + case err == nil: + successes++ + case errors.Is(err, ErrUDPDestinationCapacity): + default: + t.Fatalf("concurrent write error = %v", err) + } + } + if successes != maxUDPAllowedDestinations { + t.Fatalf("successful new destinations = %d, want %d", successes, maxUDPAllowedDestinations) + } + conn.allowedMu.RLock() + allowedCount := len(conn.allowed) + conn.allowedMu.RUnlock() + if allowedCount != maxUDPAllowedDestinations { + t.Fatalf("allowed destination count = %d, want %d", allowedCount, maxUDPAllowedDestinations) + } +} + +func TestDialContextCancellationClosesLateConnection(t *testing.T) { + core := newDelayedCore() + client := &Client{core: core} + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := client.DialContext(ctx, "tcp", "1.1.1.1:443") + result <- err + }() + <-core.started + cancel() + if err := <-result; !errors.Is(err, context.Canceled) { + t.Fatalf("DialContext error = %v, want context canceled", err) + } + close(core.release) + select { + case <-core.returned.closed: + case <-time.After(time.Second): + t.Fatal("connection returned after cancellation was not closed") + } +} + +func TestConcurrentCanceledDialsShareOneConnectionAttempt(t *testing.T) { + started := make(chan struct{}) + release := make(chan struct{}) + var startOnce sync.Once + var attempts atomic.Int64 + core := &instantCore{} + client := &Client{ + connectFunc: func() (hyclient.Client, *hyclient.HandshakeInfo, error) { + attempts.Add(1) + startOnce.Do(func() { close(started) }) + <-release + return core, &hyclient.HandshakeInfo{UDPEnabled: true}, nil + }, + } + + const callers = 128 + results := make(chan error, callers) + var callersDone sync.WaitGroup + for range callers { + callersDone.Add(1) + go func() { + defer callersDone.Done() + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + _, err := client.DialContext(ctx, "tcp", "1.1.1.1:443") + results <- err + }() + } + <-started + callersDone.Wait() + close(results) + for err := range results { + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("canceled dial error = %v, want deadline exceeded", err) + } + } + if got := attempts.Load(); got != 1 { + t.Fatalf("connection attempts = %d, want 1", got) + } + + close(release) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := client.DialContext(ctx, "tcp", "1.1.1.1:443") + if err != nil { + t.Fatal(err) + } + _ = conn.Close() + if got := attempts.Load(); got != 1 { + t.Fatalf("successful reuse started %d connection attempts, want 1", got) + } + if err := client.Close(); err != nil { + t.Fatal(err) + } +} + +func TestCanceledTCPDialsKeepUnderlyingWorkersBounded(t *testing.T) { + core := newBlockingOpenCore() + client := &Client{ + core: core, + closeCh: make(chan struct{}), + openSlots: make(chan struct{}, 2), + } + t.Cleanup(func() { _ = client.Close() }) + + for index := 1; index <= 2; index++ { + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + _, err := client.DialContext(ctx, "tcp", "1.1.1.1:443") + cancel() + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("dial %d error = %v, want deadline", index, err) + } + } + if got := core.calls.Load(); got != 2 { + t.Fatalf("underlying open calls = %d, want 2", got) + } + + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + _, err := client.DialContext(ctx, "tcp", "1.1.1.1:443") + cancel() + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("capacity-waiting dial error = %v, want deadline", err) + } + if got := core.calls.Load(); got != 2 { + t.Fatalf("capacity gate allowed %d underlying opens, want 2", got) + } + + if err := client.Close(); err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(time.Second) + for len(client.openSlots) != 0 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if got := len(client.openSlots); got != 0 { + t.Fatalf("worker slots after Close = %d, want 0", got) + } +} + +func TestAutoFallsBackAfterPrimaryUDPPathTimeout(t *testing.T) { + core := newDelayedCore() + primary := &Client{core: core} + fallback := &recordingFallback{} + var observed atomic.Int64 + auto, err := NewAutoClient(AutoConfig{ + Primary: primary, Fallback: fallback, + AttemptTimeout: 30 * time.Millisecond, + Cooldown: time.Second, + OnFallback: func(error) { + observed.Add(1) + }, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = auto.Close() }) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := auto.DialContext(ctx, "tcp", "1.1.1.1:443") + if err != nil { + t.Fatal(err) + } + _ = conn.Close() + if got := fallback.calls.Load(); got != 1 { + t.Fatalf("TLS fallback calls = %d, want 1", got) + } + if got := observed.Load(); got != 1 { + t.Fatalf("fallback observations = %d, want 1", got) + } + if got := auto.AccelerationMode(); got != "tls-fallback" { + t.Fatalf("fallback acceleration mode = %q", got) + } + if got := auto.NegotiatedTx(); got != 0 { + t.Fatalf("fallback negotiated rate = %d, want 0", got) + } + close(core.release) + select { + case <-core.returned.closed: + case <-time.After(time.Second): + t.Fatal("timed-out primary returned a connection that was not closed") + } + if _, err := auto.DialContext(ctx, "tcp", "1.1.1.1:443"); err != nil { + t.Fatalf("circuit fallback: %v", err) + } + if got := fallback.calls.Load(); got != 2 { + t.Fatalf("TLS fallback calls in cooldown = %d, want 2", got) + } + if got := observed.Load(); got != 1 { + t.Fatalf("cooldown emitted %d fallback observations, want one", got) + } +} + +func TestAdmissionControllerEnforcesMaximumAndReleasesSlot(t *testing.T) { + controller := newAdmissionController(testToken, 1) + address := netip.MustParseAddrPort("127.0.0.1:12345") + firstOK, firstID := controller.Authenticate(net.UDPAddrFromAddrPort(address), testToken, 0) + if !firstOK || firstID == "" { + t.Fatal("first authenticated connection was rejected") + } + if ok, _ := controller.Authenticate(net.UDPAddrFromAddrPort(address), testToken, 0); ok { + t.Fatal("connection beyond configured maximum was admitted") + } + if ok, _ := controller.Authenticate(net.UDPAddrFromAddrPort(address), wrongTestToken, 0); ok { + t.Fatal("wrong token was admitted") + } + controller.Disconnect(net.UDPAddrFromAddrPort(address), firstID, nil) + if ok, id := controller.Authenticate(net.UDPAddrFromAddrPort(address), testToken, 0); !ok || id == "" { + t.Fatal("released admission slot was not reusable") + } +} + +func TestClientCloseIsConcurrentIdempotentAndReturnsFirstError(t *testing.T) { + want := errors.New("close failure") + core := &closeErrorCore{err: want} + client := &Client{core: core} + const callers = 16 + errorsSeen := make(chan error, callers) + var callersDone sync.WaitGroup + for range callers { + callersDone.Add(1) + go func() { + defer callersDone.Done() + errorsSeen <- client.Close() + }() + } + callersDone.Wait() + close(errorsSeen) + for err := range errorsSeen { + if !errors.Is(err, want) { + t.Fatalf("Close error = %v, want %v", err, want) + } + } + if got := core.calls.Load(); got != 1 { + t.Fatalf("core Close calls = %d, want 1", got) + } + if _, err := client.DialContext(context.Background(), "tcp", "1.1.1.1:443"); !errors.Is(err, net.ErrClosed) { + t.Fatalf("TCP dial after Close = %v, want net.ErrClosed", err) + } + if _, err := client.DialPacket(context.Background()); !errors.Is(err, net.ErrClosed) { + t.Fatalf("UDP dial after Close = %v, want net.ErrClosed", err) + } +} + +func TestServerServeAndCloseAreOneShotAndConcurrentSafe(t *testing.T) { + core := newBlockingServerCore() + server := &Server{core: core, address: &net.UDPAddr{}, admission: newAdmissionController(testToken, 1)} + ctx, cancel := context.WithCancelCause(context.Background()) + done := make(chan error, 1) + go func() { done <- server.Serve(ctx) }() + select { + case <-core.started: + case <-time.After(time.Second): + t.Fatal("core Serve did not start") + } + if err := server.Serve(context.Background()); err == nil || !strings.Contains(err.Error(), "already serving") { + t.Fatalf("second Serve error = %v", err) + } + cancel(context.Canceled) + select { + case err := <-done: + if err != nil { + t.Fatalf("Serve cancellation error = %v", err) + } + case <-time.After(time.Second): + t.Fatal("Serve did not stop after context cancellation") + } + + const closers = 16 + var closeDone sync.WaitGroup + for range closers { + closeDone.Add(1) + go func() { + defer closeDone.Done() + if err := server.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }() + } + closeDone.Wait() + if got := core.closeCalls.Load(); got != 1 { + t.Fatalf("core Close calls = %d, want 1", got) + } +} + +func TestAutoRejectsDialsAfterClose(t *testing.T) { + core := newDelayedCore() + close(core.release) + primary := &Client{core: core} + fallback := &recordingFallback{} + auto, err := NewAutoClient(AutoConfig{ + Primary: primary, Fallback: fallback, + AttemptTimeout: time.Second, + Cooldown: time.Second, + }) + if err != nil { + t.Fatal(err) + } + if err := auto.Close(); err != nil { + t.Fatal(err) + } + if _, err := auto.DialContext(context.Background(), "tcp", "1.1.1.1:443"); !errors.Is(err, net.ErrClosed) { + t.Fatalf("TCP dial after Close = %v, want net.ErrClosed", err) + } + if _, err := auto.DialPacket(context.Background()); !errors.Is(err, net.ErrClosed) { + t.Fatalf("UDP dial after Close = %v, want net.ErrClosed", err) + } + if got := fallback.calls.Load(); got != 0 { + t.Fatalf("closed AutoClient invoked fallback %d times", got) + } +} + +func TestAutoCloseIsConcurrentIdempotentAndPreservesErrors(t *testing.T) { + primaryErr := errors.New("primary close failure") + fallbackErr := errors.New("fallback close failure") + primaryCore := &closeErrorCore{err: primaryErr} + primary := &Client{core: primaryCore} + fallback := &recordingFallback{closeErr: fallbackErr} + auto, err := NewAutoClient(AutoConfig{ + Primary: primary, Fallback: fallback, + AttemptTimeout: time.Second, + Cooldown: time.Second, + }) + if err != nil { + t.Fatal(err) + } + + const callers = 16 + errorsSeen := make(chan error, callers) + var callersDone sync.WaitGroup + for range callers { + callersDone.Add(1) + go func() { + defer callersDone.Done() + errorsSeen <- auto.Close() + }() + } + callersDone.Wait() + close(errorsSeen) + for closeErr := range errorsSeen { + if !errors.Is(closeErr, primaryErr) || !errors.Is(closeErr, fallbackErr) { + t.Fatalf("Close error = %v, want both close failures", closeErr) + } + } + if got := primaryCore.calls.Load(); got != 1 { + t.Fatalf("primary Close calls = %d, want 1", got) + } + if got := fallback.closeCalls.Load(); got != 1 { + t.Fatalf("fallback Close calls = %d, want 1", got) + } +} + +func TestAutoCloseCancelsActiveFallbackBeforeClosingIt(t *testing.T) { + primary := &Client{core: &instantCore{}} + fallback := newContextBlockingFallback() + auto, err := NewAutoClient(AutoConfig{ + Primary: primary, Fallback: fallback, + AttemptTimeout: time.Second, + Cooldown: time.Second, + }) + if err != nil { + t.Fatal(err) + } + + dialDone := make(chan error, 1) + go func() { + _, err := auto.dialFallback(context.Background(), "tcp", "1.1.1.1:443") + dialDone <- err + }() + select { + case <-fallback.started: + case <-time.After(time.Second): + t.Fatal("fallback dial did not start") + } + if err := auto.Close(); err != nil { + t.Fatal(err) + } + select { + case err := <-dialDone: + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("active fallback error = %v, want closed", err) + } + case <-time.After(time.Second): + t.Fatal("Close did not cancel the active fallback") + } + if got := fallback.closeCalls.Load(); got != 1 { + t.Fatalf("fallback close calls = %d, want 1", got) + } + if _, err := auto.dialFallback(context.Background(), "tcp", "1.1.1.1:443"); !errors.Is(err, net.ErrClosed) { + t.Fatalf("fallback dispatch after Close = %v, want closed", err) + } + if got := fallback.calls.Load(); got != 1 { + t.Fatalf("fallback was dispatched %d times, want one pre-Close call", got) + } +} + +func TestClientAndServerRejectCoreInvalidTimeoutsEagerly(t *testing.T) { + _, clientTLS := testTLSConfigs(t) + for name, modify := range map[string]func(*ClientConfig){ + "idle too short": func(config *ClientConfig) { config.MaxIdleTimeout = time.Second }, + "keepalive too short": func(config *ClientConfig) { config.KeepAlivePeriod = time.Second }, + "idle too long": func(config *ClientConfig) { config.MaxIdleTimeout = 121 * time.Second }, + } { + t.Run("client "+name, func(t *testing.T) { + config := ClientConfig{ServerAddress: "127.0.0.1:443", Token: testToken, TLSConfig: clientTLS} + modify(&config) + if _, err := NewClient(config); err == nil { + t.Fatal("invalid core timeout passed eager validation") + } + }) + } + serverTLS, _ := testTLSConfigs(t) + for name, modify := range map[string]func(*ServerConfig){ + "idle too short": func(config *ServerConfig) { config.MaxIdleTimeout = time.Second }, + "UDP idle too short": func(config *ServerConfig) { config.UDPIdleTimeout = time.Second }, + "UDP idle too long": func(config *ServerConfig) { config.UDPIdleTimeout = 601 * time.Second }, + "authentication too long": func(config *ServerConfig) { config.AuthenticationTimeout = 61 * time.Second }, + "too few unidirectional": func(config *ServerConfig) { config.MaxIncomingUniStreams = 2 }, + "too many unidirectional": func(config *ServerConfig) { config.MaxIncomingUniStreams = 1025 }, + "source connections over global": func(config *ServerConfig) { + config.MaxConnections, config.MaxClientConnections = 1, 2 + }, + "source TCP over global": func(config *ServerConfig) { + config.MaxOutboundTCP, config.MaxClientTCPHandlers = 1, 2 + }, + "source UDP over global": func(config *ServerConfig) { + config.MaxOutboundUDP, config.MaxClientUDPSessions = 1, 2 + }, + "unsafe Brutal opt-in": func(config *ServerConfig) { config.AllowClientBandwidth = true }, + } { + t.Run("server "+name, func(t *testing.T) { + config := ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Outbound: &testOutbound{}, + } + modify(&config) + if _, err := Listen(config); err == nil { + t.Fatal("invalid core timeout passed eager validation") + } + }) + } +} + +type testOutbound struct { + tcpDials atomic.Int64 + udpChecks atomic.Int64 +} + +type hyRotatingResolver struct { + answers [][]netip.Addr + calls atomic.Uint64 +} + +func (r *hyRotatingResolver) LookupNetIP(context.Context, string, string) ([]netip.Addr, error) { + call := r.calls.Add(1) + index := int(call - 1) + if index >= len(r.answers) { + index = len(r.answers) - 1 + } + return append([]netip.Addr(nil), r.answers[index]...), nil +} + +func (o *testOutbound) TCP(address string) (net.Conn, error) { + o.tcpDials.Add(1) + return net.DialTimeout("tcp", address, 2*time.Second) +} + +func (o *testOutbound) UDP(string) (hyserver.UDPConn, error) { + o.udpChecks.Add(1) + conn, err := net.ListenUDP("udp", nil) + if err != nil { + return nil, err + } + return &testUDPConn{UDPConn: conn}, nil +} + +func (o *testOutbound) CheckUDP(string) error { + o.udpChecks.Add(1) + return nil +} + +type testUDPConn struct{ *net.UDPConn } + +func (c *testUDPConn) ReadFrom(payload []byte) (int, string, error) { + n, address, err := c.ReadFromUDPAddrPort(payload) + if err != nil { + return n, "", err + } + address = netip.AddrPortFrom(address.Addr().Unmap(), address.Port()) + return n, address.String(), nil +} + +func (c *testUDPConn) WriteTo(payload []byte, address string) (int, error) { + target, err := netip.ParseAddrPort(address) + if err != nil { + return 0, err + } + return c.WriteToUDPAddrPort(payload, target) +} + +func startTestServer(t *testing.T, modify func(*ServerConfig)) (*Server, *tls.Config, *testOutbound) { + t.Helper() + serverTLS, clientTLS := testTLSConfigs(t) + outbound := &testOutbound{} + config := ServerConfig{ + Address: "127.0.0.1:0", + Token: testToken, + TLSConfig: serverTLS, + Outbound: outbound, + } + if modify != nil { + modify(&config) + } + server, err := Listen(config) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- server.Serve(ctx) }() + t.Cleanup(func() { + cancel() + _ = server.Close() + select { + case err := <-done: + if err != nil { + t.Errorf("Serve: %v", err) + } + case <-time.After(2 * time.Second): + t.Error("server did not stop") + } + }) + return server, clientTLS, outbound +} + +func newTestClient(t *testing.T, server *Server, tlsConfig *tls.Config, modify func(*ClientConfig)) *Client { + t.Helper() + config := ClientConfig{ + ServerAddress: server.Addr().String(), + Token: testToken, + TLSConfig: tlsConfig, + } + if modify != nil { + modify(&config) + } + client, err := NewClient(config) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + return client +} + +func exchangeTCP(t *testing.T, dialer transport.Dialer, address, message string) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + conn, err := dialer.DialContext(ctx, "tcp", address) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + if err := conn.SetDeadline(time.Now().Add(3 * time.Second)); err != nil { + t.Fatal(err) + } + if _, err := io.WriteString(conn, message); err != nil { + t.Fatal(err) + } + response := make([]byte, len(message)) + if _, err := io.ReadFull(conn, response); err != nil { + t.Fatal(err) + } + if string(response) != message { + t.Fatalf("response = %q, want %q", response, message) + } +} + +func startTCPEcho(t *testing.T) string { + t.Helper() + listener, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + var connections sync.WaitGroup + done := make(chan struct{}) + go func() { + defer close(done) + for { + conn, err := listener.Accept() + if err != nil { + return + } + connections.Add(1) + go func() { + defer connections.Done() + defer conn.Close() + _, _ = io.Copy(conn, conn) + }() + } + }() + t.Cleanup(func() { + _ = listener.Close() + <-done + connections.Wait() + }) + return listener.Addr().String() +} + +func startUDPEcho(t *testing.T) string { + t.Helper() + listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + go func() { + defer close(done) + buffer := make([]byte, 64<<10) + for { + n, address, err := listener.ReadFromUDPAddrPort(buffer) + if err != nil { + return + } + _, _ = listener.WriteToUDPAddrPort(buffer[:n], address) + } + }() + t.Cleanup(func() { + _ = listener.Close() + <-done + }) + return listener.LocalAddr().String() +} + +func testTLSConfigs(t *testing.T) (*tls.Config, *tls.Config) { + t.Helper() + certPEM, keyPEM, err := security.GenerateSelfSignedCertificate(security.CertificateOptions{ + Hosts: []string{"127.0.0.1"}, ValidFor: time.Hour, + }) + if err != nil { + t.Fatal(err) + } + certificate, err := tls.X509KeyPair(certPEM, keyPEM) + if err != nil { + t.Fatal(err) + } + pool := x509.NewCertPool() + if !pool.AppendCertsFromPEM(certPEM) { + t.Fatal("test certificate was not accepted as a root") + } + return &tls.Config{ + MinVersion: tls.VersionTLS13, + MaxVersion: tls.VersionTLS13, + Certificates: []tls.Certificate{certificate}, + }, &tls.Config{ + MinVersion: tls.VersionTLS13, + MaxVersion: tls.VersionTLS13, + ServerName: "127.0.0.1", + RootCAs: pool, + } +} + +type delayedCore struct { + started chan struct{} + release chan struct{} + returned *closeTrackingConn + start sync.Once + close sync.Once +} + +type blockingOpenCore struct { + release chan struct{} + close sync.Once + calls atomic.Int64 +} + +func newBlockingOpenCore() *blockingOpenCore { + return &blockingOpenCore{release: make(chan struct{})} +} + +func (c *blockingOpenCore) TCP(string) (net.Conn, error) { + c.calls.Add(1) + <-c.release + return nil, net.ErrClosed +} + +func (c *blockingOpenCore) UDP() (hyclient.HyUDPConn, error) { + return nil, errors.New("not implemented") +} + +func (c *blockingOpenCore) Close() error { + c.close.Do(func() { close(c.release) }) + return nil +} + +type closeErrorCore struct { + err error + calls atomic.Int64 +} + +type instantCore struct { + closed atomic.Bool +} + +func (c *instantCore) TCP(string) (net.Conn, error) { + if c.closed.Load() { + return nil, net.ErrClosed + } + local, peer := net.Pipe() + _ = peer.Close() + return local, nil +} + +func (c *instantCore) UDP() (hyclient.HyUDPConn, error) { + return nil, errors.New("not implemented") +} + +func (c *instantCore) Close() error { + c.closed.Store(true) + return nil +} + +func (c *closeErrorCore) TCP(string) (net.Conn, error) { return nil, net.ErrClosed } +func (c *closeErrorCore) UDP() (hyclient.HyUDPConn, error) { return nil, net.ErrClosed } +func (c *closeErrorCore) Close() error { + c.calls.Add(1) + return c.err +} + +type blockingServerCore struct { + started chan struct{} + closed chan struct{} + start sync.Once + close sync.Once + closeCalls atomic.Int64 +} + +func newBlockingServerCore() *blockingServerCore { + return &blockingServerCore{started: make(chan struct{}), closed: make(chan struct{})} +} + +func (c *blockingServerCore) Serve() error { + c.start.Do(func() { close(c.started) }) + <-c.closed + return net.ErrClosed +} + +func (c *blockingServerCore) Close() error { + c.closeCalls.Add(1) + c.close.Do(func() { close(c.closed) }) + return nil +} + +func newDelayedCore() *delayedCore { + local, peer := net.Pipe() + _ = peer.Close() + return &delayedCore{ + started: make(chan struct{}), release: make(chan struct{}), + returned: &closeTrackingConn{Conn: local, closed: make(chan struct{})}, + } +} + +func (c *delayedCore) TCP(string) (net.Conn, error) { + c.start.Do(func() { close(c.started) }) + <-c.release + return c.returned, nil +} + +func (c *delayedCore) UDP() (hyclient.HyUDPConn, error) { + return nil, errors.New("not implemented") +} + +func (c *delayedCore) Close() error { + c.close.Do(func() { _ = c.returned.Close() }) + return nil +} + +type closeTrackingConn struct { + net.Conn + closed chan struct{} + once sync.Once +} + +func (c *closeTrackingConn) Close() error { + err := c.Conn.Close() + c.once.Do(func() { close(c.closed) }) + return err +} + +type recordingFallback struct { + calls atomic.Int64 + closeCalls atomic.Int64 + closed atomic.Bool + closeErr error +} + +type contextBlockingFallback struct { + started chan struct{} + start sync.Once + calls atomic.Int64 + closeCalls atomic.Int64 +} + +func newContextBlockingFallback() *contextBlockingFallback { + return &contextBlockingFallback{started: make(chan struct{})} +} + +func (d *contextBlockingFallback) DialContext(ctx context.Context, _, _ string) (net.Conn, error) { + d.calls.Add(1) + d.start.Do(func() { close(d.started) }) + <-ctx.Done() + return nil, context.Cause(ctx) +} + +func (d *contextBlockingFallback) Close() error { + d.closeCalls.Add(1) + return nil +} + +func (d *recordingFallback) DialContext(context.Context, string, string) (net.Conn, error) { + if d.closed.Load() { + return nil, net.ErrClosed + } + d.calls.Add(1) + local, peer := net.Pipe() + _ = peer.Close() + return local, nil +} + +func (d *recordingFallback) Close() error { + d.closeCalls.Add(1) + d.closed.Store(true) + return d.closeErr +} + +var _ hyserver.Outbound = (*testOutbound)(nil) +var _ hyserver.UDPConn = (*testUDPConn)(nil) +var _ hyclient.Client = (*delayedCore)(nil) +var _ hyclient.Client = (*blockingOpenCore)(nil) +var _ hyclient.Client = (*closeErrorCore)(nil) +var _ hyclient.Client = (*instantCore)(nil) +var _ hyserver.Server = (*blockingServerCore)(nil) diff --git a/internal/hy2/server.go b/internal/hy2/server.go new file mode 100644 index 0000000..cf3dfa2 --- /dev/null +++ b/internal/hy2/server.go @@ -0,0 +1,622 @@ +package hy2 + +import ( + "context" + "crypto/tls" + "crypto/x509" + "errors" + "fmt" + "html/template" + "log/slog" + "net" + "net/http" + "net/netip" + "strings" + "sync" + "sync/atomic" + "time" + + hyserver "github.com/apernet/hysteria/core/v2/server" + "github.com/apernet/hysteria/extras/v2/obfs" + "github.com/cppla/autocar/internal/protocol" + "github.com/cppla/autocar/internal/security" +) + +const ( + defaultMaxConnections = 256 + defaultMaxClientConnections = 32 + defaultMaxStreams = 1024 + defaultMaxUniStreams = 8 + defaultMaxHTTPHeaderBytes = 16 << 10 + defaultMaxOutboundTCP = 1024 + defaultMaxOutboundUDP = 256 + defaultMaxClientTCP = 128 + defaultMaxClientUDP = 64 + maxUDPAllowedDestinations = 256 + defaultDialTimeout = 4 * time.Second + defaultRequestTimeout = 10 * time.Second + defaultUDPIdleTimeout = 60 * time.Second +) + +// ServerConfig configures the accelerated HTTP/3 relay. +type ServerConfig struct { + Address string + Token string + TLSConfig *tls.Config + Dialer *security.SafeDialer + // Outbound is an optional advanced/test adapter. Production callers should + // leave it nil so every destination is enforced by Dialer. + Outbound hyserver.Outbound + + Congestion string + BBRProfile string + MaxTx uint64 + MaxRx uint64 + + AllowClientBandwidth bool + DisableLossCompensation bool + DisableUDP bool + DisablePathMTUDiscovery bool + DisableGSO bool + ObfuscationKey []byte + MaxIdleTimeout time.Duration + UDPIdleTimeout time.Duration + DialTimeout time.Duration + MaxConcurrentStreams int + MaxIncomingUniStreams int + MaxConnections int + MaxClientConnections int + MaxOutboundTCP int + MaxOutboundUDP int + MaxClientTCPHandlers int + MaxClientUDPSessions int + TCPRequestTimeout time.Duration + AuthenticationTimeout time.Duration + MasqueradeHandler http.Handler +} + +// Server is an authenticated Hysteria v2 relay. +type Server struct { + core hyserver.Server + address net.Addr + admission *admissionController + + serveMu sync.Mutex + serving bool + closed atomic.Bool + close sync.Once + closeErr error +} + +// Listen binds the UDP socket and constructs the HTTP/3 relay. +func Listen(config ServerConfig) (*Server, error) { + if err := validateServerConfig(config); err != nil { + return nil, err + } + config.Congestion = normalizeCongestion(config.Congestion) + config.BBRProfile = normalizeBBRProfile(config.BBRProfile) + if config.MaxConnections == 0 { + config.MaxConnections = defaultMaxConnections + } + if config.MaxClientConnections == 0 { + config.MaxClientConnections = min(defaultMaxClientConnections, config.MaxConnections) + } + if config.MaxConcurrentStreams == 0 { + config.MaxConcurrentStreams = defaultMaxStreams + } + if config.MaxIncomingUniStreams == 0 { + config.MaxIncomingUniStreams = defaultMaxUniStreams + } + if config.MaxOutboundTCP == 0 { + config.MaxOutboundTCP = defaultMaxOutboundTCP + } + if config.MaxOutboundUDP == 0 { + config.MaxOutboundUDP = defaultMaxOutboundUDP + } + if config.MaxClientTCPHandlers == 0 { + config.MaxClientTCPHandlers = min(defaultMaxClientTCP, config.MaxOutboundTCP) + } + if config.MaxClientUDPSessions == 0 { + config.MaxClientUDPSessions = min(defaultMaxClientUDP, config.MaxOutboundUDP) + } + if config.DialTimeout == 0 { + config.DialTimeout = defaultDialTimeout + } + if config.TCPRequestTimeout == 0 { + config.TCPRequestTimeout = defaultRequestTimeout + } + if config.AuthenticationTimeout == 0 { + config.AuthenticationTimeout = defaultRequestTimeout + } + if config.UDPIdleTimeout == 0 { + config.UDPIdleTimeout = defaultUDPIdleTimeout + } + + packetConn, err := net.ListenPacket("udp", config.Address) + if err != nil { + return nil, fmt.Errorf("hy2: listen UDP: %w", err) + } + address := packetConn.LocalAddr() + if len(config.ObfuscationKey) != 0 { + wrapped, wrapErr := obfs.WrapPacketConnSalamander(packetConn, config.ObfuscationKey) + if wrapErr != nil { + _ = packetConn.Close() + return nil, fmt.Errorf("hy2: enable Salamander: %w", wrapErr) + } + packetConn = wrapped + } + + admission := newAdmissionController(config.Token, config.MaxConnections) + masquerade := config.MasqueradeHandler + if masquerade == nil { + masquerade = NewCoverHandler("") + } + tlsConfig := config.TLSConfig.Clone() + outbound := config.Outbound + if outbound == nil { + outbound = &safeOutbound{ + dialer: config.Dialer, + timeout: config.DialTimeout, + tcpSlots: make(chan struct{}, config.MaxOutboundTCP), + udpSlots: make(chan struct{}, config.MaxOutboundUDP), + } + } + core, err := hyserver.NewServer(&hyserver.Config{ + Conn: packetConn, + TLSConfig: hyserver.TLSConfig{ + Certificates: append([]tls.Certificate(nil), tlsConfig.Certificates...), + GetCertificate: tlsConfig.GetCertificate, + ClientCAs: cloneCertPool(tlsConfig.ClientCAs), + ECHKeys: append([]tls.EncryptedClientHelloKey(nil), tlsConfig.EncryptedClientHelloKeys...), + GetECHKeys: tlsConfig.GetEncryptedClientHelloKeys, + }, + QUICConfig: hyserver.QUICConfig{ + MaxIdleTimeout: config.MaxIdleTimeout, + MaxIncomingStreams: int64(config.MaxConcurrentStreams), + MaxIncomingUniStreams: int64(config.MaxIncomingUniStreams), + DisablePathMTUDiscovery: config.DisablePathMTUDiscovery, + DisableGSO: config.DisableGSO, + }, + Outbound: outbound, + CongestionConfig: hyserver.CongestionConfig{ + Type: config.Congestion, + BBRProfile: config.BBRProfile, + }, + BandwidthConfig: hyserver.BandwidthConfig{ + MaxTx: config.MaxTx, + MaxRx: config.MaxRx, + DisableLossCompensation: config.DisableLossCompensation, + }, + IgnoreClientBandwidth: !config.AllowClientBandwidth, + DisableUDP: config.DisableUDP, + UDPIdleTimeout: config.UDPIdleTimeout, + MaxConnections: config.MaxConnections, + MaxClientConnections: config.MaxClientConnections, + MaxTCPHandlers: config.MaxOutboundTCP, + MaxClientTCPHandlers: config.MaxClientTCPHandlers, + TCPRequestTimeout: config.TCPRequestTimeout, + AuthenticationTimeout: config.AuthenticationTimeout, + MaxHTTPHeaderBytes: defaultMaxHTTPHeaderBytes, + MaxUDPSessions: config.MaxOutboundUDP, + MaxClientUDPSessions: config.MaxClientUDPSessions, + Authenticator: admission, + EventLogger: admission, + MasqHandler: masquerade, + }) + if err != nil { + _ = packetConn.Close() + return nil, fmt.Errorf("hy2: create server: %w", err) + } + return &Server{core: core, address: address, admission: admission}, nil +} + +func validateServerConfig(config ServerConfig) error { + if config.Address == "" { + return errors.New("hy2: listen address is required") + } + if len(config.Token) < protocol.MinTokenLength || len(config.Token) > protocol.MaxTokenLength { + return fmt.Errorf("hy2: token length must be between %d and %d bytes", protocol.MinTokenLength, protocol.MaxTokenLength) + } + if config.TLSConfig == nil { + return errors.New("hy2: server TLS config is required") + } + if err := validateServerTLSPolicy(config.TLSConfig); err != nil { + return err + } + if len(config.TLSConfig.Certificates) == 0 && config.TLSConfig.GetCertificate == nil { + return errors.New("hy2: server TLS certificate is required") + } + if config.Dialer == nil && config.Outbound == nil { + return errors.New("hy2: safe outbound dialer is required") + } + congestion := normalizeCongestion(config.Congestion) + if congestion != CongestionBBR && congestion != CongestionReno { + return fmt.Errorf("hy2: unsupported congestion controller %q", config.Congestion) + } + profile := normalizeBBRProfile(config.BBRProfile) + if congestion == CongestionBBR && profile != BBRConservative && profile != BBRStandard && profile != BBRAggressive { + return fmt.Errorf("hy2: unsupported BBR profile %q", config.BBRProfile) + } + for name, bandwidth := range map[string]uint64{"MaxTx": config.MaxTx, "MaxRx": config.MaxRx} { + if bandwidth != 0 && bandwidth < minimumBandwidth { + return fmt.Errorf("hy2: %s must be zero or at least %d bytes/s", name, minimumBandwidth) + } + if bandwidth > maximumBandwidth { + return fmt.Errorf("hy2: %s must not exceed %d bytes/s", name, maximumBandwidth) + } + } + if config.AllowClientBandwidth && (config.MaxTx == 0 || config.MaxRx == 0) { + return errors.New("hy2: allowing client bandwidth requires finite MaxTx and MaxRx ceilings") + } + if len(config.ObfuscationKey) != 0 && len(config.ObfuscationKey) < minimumObfsKey { + return fmt.Errorf("hy2: obfuscation key must be at least %d bytes", minimumObfsKey) + } + if config.MaxConnections < 0 || config.MaxConnections > 65536 { + return errors.New("hy2: maximum connections must be zero or at most 65536") + } + if config.MaxClientConnections < 0 || config.MaxClientConnections > 65536 { + return errors.New("hy2: per-source connection limit must be zero or at most 65536") + } + if config.MaxClientConnections > 0 && config.MaxConnections > 0 && config.MaxClientConnections > config.MaxConnections { + return errors.New("hy2: per-source connection limit cannot exceed the global connection limit") + } + if config.MaxConcurrentStreams < 0 || config.MaxConcurrentStreams > 65536 || (config.MaxConcurrentStreams > 0 && config.MaxConcurrentStreams < 8) { + return errors.New("hy2: maximum streams must be zero or between 8 and 65536") + } + if config.MaxIncomingUniStreams < 0 || config.MaxIncomingUniStreams > 1024 || (config.MaxIncomingUniStreams > 0 && config.MaxIncomingUniStreams < 3) { + return errors.New("hy2: maximum unidirectional streams must be zero or between 3 and 1024") + } + if config.MaxOutboundTCP < 0 || config.MaxOutboundTCP > 65536 || config.MaxOutboundUDP < 0 || config.MaxOutboundUDP > 65536 { + return errors.New("hy2: outbound connection limits must be zero or at most 65536") + } + if config.MaxClientUDPSessions < 0 || config.MaxClientUDPSessions > 65536 { + return errors.New("hy2: per-source UDP session limit must be zero or at most 65536") + } + if config.MaxClientTCPHandlers < 0 || config.MaxClientTCPHandlers > 65536 { + return errors.New("hy2: per-source TCP handler limit must be zero or at most 65536") + } + if config.MaxClientTCPHandlers > 0 && config.MaxOutboundTCP > 0 && config.MaxClientTCPHandlers > config.MaxOutboundTCP { + return errors.New("hy2: per-source TCP handler limit cannot exceed the global TCP limit") + } + if config.MaxClientUDPSessions > 0 && config.MaxOutboundUDP > 0 && config.MaxClientUDPSessions > config.MaxOutboundUDP { + return errors.New("hy2: per-source UDP session limit cannot exceed the global UDP limit") + } + if config.DialTimeout < 0 { + return errors.New("hy2: outbound dial timeout cannot be negative") + } + if config.TCPRequestTimeout != 0 && (config.TCPRequestTimeout < time.Second || config.TCPRequestTimeout > 60*time.Second) { + return errors.New("hy2: TCP request timeout must be zero or between 1s and 60s") + } + if config.AuthenticationTimeout != 0 && (config.AuthenticationTimeout < time.Second || config.AuthenticationTimeout > 60*time.Second) { + return errors.New("hy2: authentication timeout must be zero or between 1s and 60s") + } + if config.MaxIdleTimeout != 0 && (config.MaxIdleTimeout < 4*time.Second || config.MaxIdleTimeout > 120*time.Second) { + return errors.New("hy2: maximum idle timeout must be zero or between 4s and 120s") + } + if config.UDPIdleTimeout != 0 && (config.UDPIdleTimeout < 2*time.Second || config.UDPIdleTimeout > 600*time.Second) { + return errors.New("hy2: UDP idle timeout must be zero or between 2s and 600s") + } + return nil +} + +func cloneCertPool(pool *x509.CertPool) *x509.CertPool { + if pool == nil { + return nil + } + return pool.Clone() +} + +// Addr returns the bound UDP address. +func (s *Server) Addr() net.Addr { return s.address } + +// Serve runs until ctx is canceled, Close is called or the listener fails. +func (s *Server) Serve(ctx context.Context) error { + if ctx == nil { + return errors.New("hy2: nil serve context") + } + s.serveMu.Lock() + if s.serving { + s.serveMu.Unlock() + return errors.New("hy2: server already serving") + } + s.serving = true + s.serveMu.Unlock() + + done := make(chan error, 1) + go func() { done <- s.core.Serve() }() + select { + case err := <-done: + if s.closed.Load() || errors.Is(err, net.ErrClosed) { + return nil + } + return fmt.Errorf("hy2: serve: %w", err) + case <-ctx.Done(): + _ = s.Close() + <-done + if cause := context.Cause(ctx); cause != nil && !errors.Is(cause, context.Canceled) { + return cause + } + return nil + } +} + +// Close stops the listener and all active sessions. +func (s *Server) Close() error { + s.close.Do(func() { + s.closed.Store(true) + s.closeErr = s.core.Close() + }) + return s.closeErr +} + +type admissionController struct { + token string + slots chan struct{} + next atomic.Uint64 + lastTx atomic.Uint64 + open sync.Map +} + +func newAdmissionController(token string, maximum int) *admissionController { + return &admissionController{token: token, slots: make(chan struct{}, maximum)} +} + +func (a *admissionController) Authenticate(addr net.Addr, provided string, _ uint64) (bool, string) { + if !security.VerifyToken(provided, a.token) { + return false, "" + } + select { + case a.slots <- struct{}{}: + default: + return false, "" + } + id := fmt.Sprintf("%s#%d", addr.String(), a.next.Add(1)) + a.open.Store(id, struct{}{}) + return true, id +} + +func (a *admissionController) Connect(addr net.Addr, id string, tx uint64) { + a.lastTx.Store(tx) + slog.Debug("hy2 session authenticated", "remote", addr.String(), "id", id, "tx_bytes_per_second", tx) +} + +func (a *admissionController) Disconnect(addr net.Addr, id string, err error) { + if _, loaded := a.open.LoadAndDelete(id); loaded { + <-a.slots + } + slog.Debug("hy2 session disconnected", "remote", addr.String(), "id", id, "error", err) +} + +func (a *admissionController) TCPRequest(net.Addr, string, string) {} +func (a *admissionController) TCPError(net.Addr, string, string, error) {} +func (a *admissionController) UDPRequest(net.Addr, string, uint32, string) {} +func (a *admissionController) UDPError(net.Addr, string, uint32, error) {} + +type safeOutbound struct { + dialer *security.SafeDialer + timeout time.Duration + tcpSlots chan struct{} + udpSlots chan struct{} +} + +var ( + ErrOutboundCapacity = errors.New("hy2: outbound capacity exhausted") + ErrUDPDestinationCapacity = errors.New("hy2: UDP destination capacity exhausted") +) + +func (o *safeOutbound) TCP(address string) (net.Conn, error) { + if !acquireSlot(o.tcpSlots) { + return nil, fmt.Errorf("%w: TCP", ErrOutboundCapacity) + } + ctx, cancel := context.WithTimeout(context.Background(), o.timeout) + defer cancel() + conn, err := o.dialer.DialContext(ctx, "tcp", address) + if err != nil { + releaseSlot(o.tcpSlots) + return nil, err + } + return &releaseConn{Conn: conn, release: func() { releaseSlot(o.tcpSlots) }}, nil +} + +func (o *safeOutbound) UDP(address string) (hyserver.UDPConn, error) { + if !acquireSlot(o.udpSlots) { + return nil, fmt.Errorf("%w: UDP", ErrOutboundCapacity) + } + release := true + defer func() { + if release { + releaseSlot(o.udpSlots) + } + }() + // Hysteria calls UDP directly for the first datagram in a session and uses + // CheckUDP only when the destination later changes. Enforce policy here as + // well so the initial address cannot bypass the relay's SSRF boundary. + ctx, cancel := context.WithTimeout(context.Background(), o.timeout) + defer cancel() + if _, err := o.dialer.ResolveUDPContext(ctx, address); err != nil { + return nil, err + } + conn, err := net.ListenUDP("udp", nil) + if err != nil { + return nil, err + } + release = false + return &safeUDPConn{ + conn: conn, dialer: o.dialer, timeout: o.timeout, + release: func() { releaseSlot(o.udpSlots) }, + }, nil +} + +func (o *safeOutbound) CheckUDP(address string) error { + ctx, cancel := context.WithTimeout(context.Background(), o.timeout) + defer cancel() + _, err := o.dialer.ResolveUDPContext(ctx, address) + return err +} + +type safeUDPConn struct { + conn *net.UDPConn + dialer *security.SafeDialer + timeout time.Duration + allowedMu sync.RWMutex + // allowed contains only destinations that passed policy and a successful + // socket write. Entries are never evicted: removing one could cause a valid + // delayed reply to be mistaken for an unsolicited packet. + allowed map[netip.AddrPort]struct{} + release func() + close sync.Once + closeErr error +} + +type releaseConn struct { + net.Conn + release func() + close sync.Once + closeErr error +} + +func acquireSlot(slots chan struct{}) bool { + if slots == nil { + return true + } + select { + case slots <- struct{}{}: + return true + default: + return false + } +} + +func releaseSlot(slots chan struct{}) { + if slots != nil { + <-slots + } +} + +func (c *releaseConn) Close() error { + c.close.Do(func() { + c.closeErr = c.Conn.Close() + if c.release != nil { + c.release() + } + }) + return c.closeErr +} + +func (c *safeUDPConn) ReadFrom(buffer []byte) (int, string, error) { + for { + n, address, err := c.conn.ReadFromUDPAddrPort(buffer) + if err != nil { + return n, "", err + } + address = netip.AddrPortFrom(address.Addr().Unmap(), address.Port()) + if c.destinationAllowed(address) { + return n, address.String(), nil + } + } +} + +func (c *safeUDPConn) WriteTo(payload []byte, address string) (int, error) { + ctx, cancel := context.WithTimeout(context.Background(), c.timeout) + defer cancel() + addresses, err := c.dialer.ResolveUDPContext(ctx, address) + if err != nil { + return 0, err + } + var writeErrors []error + for _, candidate := range addresses { + candidate = netip.AddrPortFrom(candidate.Addr().Unmap(), candidate.Port()) + n, writeErr := c.writeToDestination(payload, candidate) + if writeErr == nil { + return n, nil + } + writeErrors = append(writeErrors, writeErr) + } + return 0, errors.Join(writeErrors...) +} + +func (c *safeUDPConn) destinationAllowed(address netip.AddrPort) bool { + c.allowedMu.RLock() + _, ok := c.allowed[address] + c.allowedMu.RUnlock() + return ok +} + +func (c *safeUDPConn) writeToDestination(payload []byte, address netip.AddrPort) (int, error) { + // Existing destinations need no admission change and remain usable even + // after the fixed-size set is full. + if c.destinationAllowed(address) { + return c.conn.WriteToUDPAddrPort(payload, address) + } + + // Serialize first writes so concurrent successful sends cannot overfill the + // set. Keep the lock through the socket write: a reply read after that write + // waits until authorization is recorded, rather than being dropped in the + // small interval between the two operations. + c.allowedMu.Lock() + defer c.allowedMu.Unlock() + if _, ok := c.allowed[address]; ok { + return c.conn.WriteToUDPAddrPort(payload, address) + } + if len(c.allowed) >= maxUDPAllowedDestinations { + return 0, ErrUDPDestinationCapacity + } + n, err := c.conn.WriteToUDPAddrPort(payload, address) + if err != nil { + return n, err + } + if c.allowed == nil { + c.allowed = make(map[netip.AddrPort]struct{}, maxUDPAllowedDestinations) + } + c.allowed[address] = struct{}{} + return n, nil +} + +func (c *safeUDPConn) Close() error { + c.close.Do(func() { + c.closeErr = c.conn.Close() + if c.release != nil { + c.release() + } + }) + return c.closeErr +} + +// CoverHandler serves a small neutral HTTP site for unauthenticated and +// ordinary HTTP/3 requests. This makes active probes observe a valid web +// service instead of an AutoCAR-specific protocol error. +type CoverHandler struct { + serverName string + page *template.Template +} + +// NewCoverHandler creates the built-in HTTP/3 cover. serverName is optional +// and is HTML-escaped by the template package. +func NewCoverHandler(serverName string) http.Handler { + serverName = strings.TrimSpace(serverName) + if serverName == "" { + serverName = "Service" + } + page := template.Must(template.New("cover").Parse("{{.}}

{{.}}

The service is online.

")) + return &CoverHandler{serverName: serverName, page: page} +} + +func (h *CoverHandler) ServeHTTP(response http.ResponseWriter, request *http.Request) { + response.Header().Set("Cache-Control", "public, max-age=300") + response.Header().Set("Content-Type", "text/html; charset=utf-8") + response.Header().Set("X-Content-Type-Options", "nosniff") + if (request.Method != http.MethodGet && request.Method != http.MethodHead) || request.URL.Path != "/" { + http.NotFound(response, request) + return + } + response.WriteHeader(http.StatusOK) + if request.Method == http.MethodHead { + return + } + _ = h.page.Execute(response, h.serverName) +} + +var _ hyserver.Authenticator = (*admissionController)(nil) +var _ hyserver.EventLogger = (*admissionController)(nil) +var _ hyserver.Outbound = (*safeOutbound)(nil) +var _ hyserver.UDPConn = (*safeUDPConn)(nil) diff --git a/internal/hy2/tls_policy.go b/internal/hy2/tls_policy.go new file mode 100644 index 0000000..b606639 --- /dev/null +++ b/internal/hy2/tls_policy.go @@ -0,0 +1,84 @@ +package hy2 + +import ( + "crypto/tls" + "errors" + "fmt" +) + +// The Hysteria core intentionally exposes a smaller TLS surface than +// crypto/tls.Config. Validate every security policy that cannot be faithfully +// copied before opening a socket. Silently dropping one of these callbacks can +// turn an application-specific certificate or resumption policy into the +// default policy. +func validateClientTLSPolicy(config *tls.Config) error { + if err := validateTLS13Versions(config, "client"); err != nil { + return err + } + switch { + case config.VerifyConnection != nil: + return errors.New("hy2: client TLS VerifyConnection is unsupported") + case config.EncryptedClientHelloRejectionVerify != nil: + return errors.New("hy2: client TLS EncryptedClientHelloRejectionVerify is unsupported") + case config.Time != nil: + return errors.New("hy2: client TLS custom Time is unsupported") + case config.Rand != nil: + return errors.New("hy2: client TLS custom Rand is unsupported") + case len(config.CurvePreferences) != 0: + return errors.New("hy2: client TLS custom CurvePreferences are unsupported") + } + return nil +} + +func validateServerTLSPolicy(config *tls.Config) error { + if err := validateTLS13Versions(config, "server"); err != nil { + return err + } + switch { + case config.GetConfigForClient != nil: + return errors.New("hy2: server TLS GetConfigForClient is unsupported") + case config.VerifyConnection != nil: + return errors.New("hy2: server TLS VerifyConnection is unsupported") + case config.VerifyPeerCertificate != nil: + return errors.New("hy2: server TLS VerifyPeerCertificate is unsupported") + case config.Time != nil: + return errors.New("hy2: server TLS custom Time is unsupported") + case config.Rand != nil: + return errors.New("hy2: server TLS custom Rand is unsupported") + case len(config.CurvePreferences) != 0: + return errors.New("hy2: server TLS custom CurvePreferences are unsupported") + case config.NameToCertificate != nil: + return errors.New("hy2: server TLS NameToCertificate is unsupported; use Certificates or GetCertificate") + case config.SessionTicketsDisabled: + return errors.New("hy2: server TLS SessionTicketsDisabled is unsupported") + case config.SessionTicketKey != ([32]byte{}): + return errors.New("hy2: server TLS custom SessionTicketKey is unsupported") + case config.WrapSession != nil: + return errors.New("hy2: server TLS WrapSession is unsupported") + case config.UnwrapSession != nil: + return errors.New("hy2: server TLS UnwrapSession is unsupported") + } + + // Hysteria exposes client CAs rather than the full ClientAuth enum. Its + // exact mapping is either no client certificate, or strict verified mTLS. + // Reject every intermediate/custom policy instead of silently strengthening + // or weakening it. + if config.ClientCAs == nil { + if config.ClientAuth != tls.NoClientCert { + return errors.New("hy2: server TLS ClientAuth requires ClientCAs and must be RequireAndVerifyClientCert") + } + } else if config.ClientAuth != tls.RequireAndVerifyClientCert { + return errors.New("hy2: server TLS ClientCAs require ClientAuth RequireAndVerifyClientCert") + } + return nil +} + +func validateTLS13Versions(config *tls.Config, role string) error { + if config.MaxVersion != 0 && config.MaxVersion < tls.VersionTLS13 { + return fmt.Errorf("hy2: %s TLS MaxVersion excludes TLS 1.3", role) + } + if config.MinVersion > tls.VersionTLS13 { + return fmt.Errorf("hy2: %s TLS MinVersion excludes TLS 1.3", role) + } + return nil +} diff --git a/internal/proxy/config.go b/internal/proxy/config.go index df3c971..5509f44 100644 --- a/internal/proxy/config.go +++ b/internal/proxy/config.go @@ -38,6 +38,7 @@ type Config struct { type serverConfig struct { dialer transport.Dialer + packetDialer transport.PacketDialer authenticator Authenticator handshakeTimeout time.Duration dialTimeout time.Duration @@ -67,8 +68,10 @@ func normalizeConfig(cfg Config) (serverConfig, error) { if cfg.MaxConnections == 0 { cfg.MaxConnections = 1024 } + packetDialer, _ := cfg.Dialer.(transport.PacketDialer) return serverConfig{ dialer: cfg.Dialer, + packetDialer: packetDialer, authenticator: cfg.Authenticator, handshakeTimeout: cfg.HandshakeTimeout, dialTimeout: cfg.DialTimeout, diff --git a/internal/proxy/socks5.go b/internal/proxy/socks5.go index 19e99e3..efd76ca 100644 --- a/internal/proxy/socks5.go +++ b/internal/proxy/socks5.go @@ -1,6 +1,7 @@ package proxy import ( + "bytes" "context" "encoding/binary" "errors" @@ -8,8 +9,12 @@ import ( "io" "net" "os" + "strconv" + "sync" "syscall" "time" + + "github.com/cppla/autocar/internal/transport" ) const ( @@ -38,10 +43,14 @@ const ( socksReplyAddressUnsupported = 0x08 userPasswordVersion = 0x01 + + maxUDPDatagramSize = 65507 ) -// SOCKS5Server is an RFC 1928 CONNECT proxy. BIND and UDP ASSOCIATE receive a -// standards-compliant "command not supported" response. +var errSOCKSUDPFragmented = errors.New("socks5: fragmented UDP datagram") + +// SOCKS5Server is an RFC 1928 CONNECT and, when the configured transport also +// implements transport.PacketDialer, UDP ASSOCIATE proxy. BIND is unsupported. type SOCKS5Server struct { cfg serverConfig lifecycle *serverLifecycle @@ -100,7 +109,15 @@ func (s *SOCKS5Server) serveConn(client net.Conn) { } return } - if request.command != socksCommandConnect { + if request.command == socksCommandBind { + _ = writeSOCKSReply(client, socksReplyCommandUnsupported, nil) + return + } + if request.command == socksCommandUDP && s.cfg.packetDialer == nil { + _ = writeSOCKSReply(client, socksReplyCommandUnsupported, nil) + return + } + if request.command != socksCommandConnect && request.command != socksCommandUDP { _ = writeSOCKSReply(client, socksReplyCommandUnsupported, nil) return } @@ -108,7 +125,14 @@ func (s *SOCKS5Server) serveConn(client net.Conn) { // parsed. Remote dialing has its own timeout and must not accidentally be // shortened by the handshake deadline. _ = client.SetDeadline(time.Time{}) + if request.command == socksCommandUDP { + s.serveUDPAssociate(client, request) + return + } + s.serveConnect(client, request) +} +func (s *SOCKS5Server) serveConnect(client net.Conn, request socksRequest) { ctx := context.Background() cancel := func() {} if s.cfg.dialTimeout > 0 { @@ -135,6 +159,356 @@ func (s *SOCKS5Server) serveConn(client net.Conn) { _ = relay(client, upstream, s.cfg.idleTimeout) } +func (s *SOCKS5Server) serveUDPAssociate(client net.Conn, request socksRequest) { + ctx := context.Background() + cancel := func() {} + if s.cfg.dialTimeout > 0 { + ctx, cancel = context.WithTimeout(ctx, s.cfg.dialTimeout) + } + defer cancel() + + peerIP, err := addressIP(client.RemoteAddr()) + if err != nil { + _ = writeSOCKSReply(client, socksReplyGeneralFailure, nil) + return + } + requestedPort, err := validateUDPAssociateRequest(ctx, request, peerIP, net.DefaultResolver.LookupIPAddr) + if err != nil { + reply := byte(socksReplyGeneralFailure) + var protocolErr *socksProtocolError + if errors.As(err, &protocolErr) { + reply = protocolErr.reply + } + _ = writeSOCKSReply(client, reply, nil) + return + } + + udpConn, err := listenSOCKSUDP(client) + if err != nil { + _ = writeSOCKSReply(client, socksReplyGeneralFailure, nil) + return + } + defer udpConn.Close() + + upstream, err := s.cfg.packetDialer.DialPacket(ctx) + if err != nil || upstream == nil { + if err == nil { + err = errors.New("socks5: packet dialer returned a nil connection") + } + _ = writeSOCKSReply(client, socksReplyForError(err), nil) + return + } + upstream = &closeOncePacketConn{PacketConn: upstream} + defer upstream.Close() + cancel() + + if s.cfg.handshakeTimeout > 0 { + _ = client.SetWriteDeadline(time.Now().Add(s.cfg.handshakeTimeout)) + } + if err := writeSOCKSReply(client, socksReplySucceeded, udpConn.LocalAddr()); err != nil { + return + } + _ = client.SetWriteDeadline(time.Time{}) + + endpoint := &socksUDPClientEndpoint{ + peerIP: peerIP, + requestedPort: requestedPort, + } + runSOCKSUDPAssociation(client, udpConn, upstream, endpoint, s.cfg.idleTimeout) +} + +type lookupIPFunc func(context.Context, string) ([]net.IPAddr, error) + +// validateUDPAssociateRequest applies the RFC 1928 meaning of DST.ADDR and +// DST.PORT: they describe the endpoint from which the client expects to send +// UDP packets. An unspecified address selects the TCP control connection's +// peer. A concrete IP, or the resolved address set for a domain, is accepted +// only when it includes that same peer. The data path independently checks the +// actual packet source and pins a zero port on the first valid datagram. +func validateUDPAssociateRequest( + ctx context.Context, + request socksRequest, + peerIP net.IP, + lookupIP lookupIPFunc, +) (int, error) { + if peerIP == nil { + return 0, &socksProtocolError{ + reply: socksReplyGeneralFailure, + err: errors.New("socks5: missing UDP control peer IP"), + } + } + + switch request.addressType { + case socksAddressIPv4, socksAddressIPv6: + requestedIP := net.ParseIP(request.host) + if requestedIP == nil { + return 0, &socksProtocolError{ + reply: socksReplyAddressUnsupported, + err: errors.New("socks5: invalid UDP associate IP"), + } + } + if !requestedIP.IsUnspecified() && !requestedIP.Equal(peerIP) { + return 0, &socksProtocolError{ + reply: socksReplyNotAllowed, + err: errors.New("socks5: UDP associate address does not match the control peer"), + } + } + case socksAddressDomain: + if lookupIP == nil { + return 0, &socksProtocolError{ + reply: socksReplyHostUnreachable, + err: errors.New("socks5: no resolver for UDP associate domain"), + } + } + addresses, err := lookupIP(ctx, request.host) + if err != nil { + return 0, &socksProtocolError{ + reply: socksReplyHostUnreachable, + err: fmt.Errorf("socks5: resolve UDP associate domain: %w", err), + } + } + matched := false + for _, address := range addresses { + if address.IP.Equal(peerIP) { + matched = true + break + } + } + if !matched { + return 0, &socksProtocolError{ + reply: socksReplyNotAllowed, + err: errors.New("socks5: UDP associate domain does not resolve to the control peer"), + } + } + default: + return 0, &socksProtocolError{ + reply: socksReplyAddressUnsupported, + err: errors.New("socks5: unsupported UDP associate address type"), + } + } + return int(request.port), nil +} + +func listenSOCKSUDP(client net.Conn) (*net.UDPConn, error) { + localIP, err := addressIP(client.LocalAddr()) + if err != nil { + return nil, err + } + if ip4 := localIP.To4(); ip4 != nil { + return net.ListenUDP("udp4", &net.UDPAddr{IP: ip4}) + } + ip16 := localIP.To16() + if ip16 == nil { + return nil, errors.New("socks5: TCP listener has no IP address") + } + return net.ListenUDP("udp6", &net.UDPAddr{IP: ip16}) +} + +func addressIP(address net.Addr) (net.IP, error) { + switch value := address.(type) { + case *net.TCPAddr: + if value.IP != nil { + return append(net.IP(nil), value.IP...), nil + } + case *net.UDPAddr: + if value.IP != nil { + return append(net.IP(nil), value.IP...), nil + } + } + if address == nil { + return nil, errors.New("socks5: missing socket address") + } + host, _, err := net.SplitHostPort(address.String()) + if err != nil { + return nil, err + } + ip := net.ParseIP(host) + if ip == nil { + return nil, errors.New("socks5: socket address is not an IP address") + } + return ip, nil +} + +type socksUDPClientEndpoint struct { + mu sync.RWMutex + peerIP net.IP + requestedPort int + address *net.UDPAddr +} + +type closeOncePacketConn struct { + transport.PacketConn + once sync.Once + err error +} + +func (c *closeOncePacketConn) Close() error { + c.once.Do(func() { c.err = c.PacketConn.Close() }) + return c.err +} + +func (c *closeOncePacketConn) MaxPayloadSize() int { + if sized, ok := c.PacketConn.(transport.PacketPayloadSizer); ok { + return sized.MaxPayloadSize() + } + return 0 +} + +// accept records the first valid source port when the UDP ASSOCIATE request +// specified port zero. Every datagram must originate from the control TCP +// connection's peer IP, preventing the relay from becoming an open UDP proxy. +func (e *socksUDPClientEndpoint) accept(address *net.UDPAddr) bool { + if address == nil || !address.IP.Equal(e.peerIP) { + return false + } + e.mu.Lock() + defer e.mu.Unlock() + if e.requestedPort != 0 && address.Port != e.requestedPort { + return false + } + if e.address == nil { + e.address = &net.UDPAddr{ + IP: append(net.IP(nil), address.IP...), + Port: address.Port, + Zone: address.Zone, + } + return true + } + return address.Port == e.address.Port && address.Zone == e.address.Zone +} + +func (e *socksUDPClientEndpoint) current() *net.UDPAddr { + e.mu.RLock() + defer e.mu.RUnlock() + if e.address == nil { + return nil + } + return &net.UDPAddr{ + IP: append(net.IP(nil), e.address.IP...), + Port: e.address.Port, + Zone: e.address.Zone, + } +} + +func runSOCKSUDPAssociation( + control net.Conn, + local *net.UDPConn, + upstream transport.PacketConn, + endpoint *socksUDPClientEndpoint, + idleTimeout time.Duration, +) { + maxPayloadSize := maxUDPDatagramSize + if sized, ok := upstream.(transport.PacketPayloadSizer); ok { + if limit := sized.MaxPayloadSize(); limit > 0 && limit < maxPayloadSize { + maxPayloadSize = limit + } + } + finished := make(chan struct{}, 3) + activity := make(chan struct{}, 1) + signalActivity := func() { + select { + case activity <- struct{}{}: + default: + } + } + finish := func() { finished <- struct{}{} } + + go func() { + defer finish() + // One extra byte lets us detect and drop oversized IPv6 UDP payloads + // instead of forwarding a silently truncated 65,507-byte prefix. + buffer := make([]byte, maxUDPDatagramSize+1) + for { + n, source, err := local.ReadFromUDP(buffer) + if err != nil { + return + } + payload, target, err := parseSOCKSUDPDatagram(buffer[:n]) + if err != nil || len(payload) > maxPayloadSize || !endpoint.accept(source) { + continue + } + if err := upstream.Send(payload, target); err != nil { + return + } + signalActivity() + } + }() + + go func() { + defer finish() + for { + payload, source, err := upstream.Receive() + if err != nil { + return + } + clientAddress := endpoint.current() + if clientAddress == nil { + // An unsolicited upstream packet cannot be routed safely before + // the client's source endpoint has been validated. + continue + } + packet, err := buildSOCKSUDPDatagram(payload, source) + if err != nil { + continue + } + if _, err := local.WriteToUDP(packet, clientAddress); err != nil { + return + } + signalActivity() + } + }() + + go func() { + defer finish() + buffer := make([]byte, 1) + for { + if _, err := control.Read(buffer); err != nil { + return + } + } + }() + + var timer *time.Timer + var idle <-chan time.Time + if idleTimeout > 0 { + timer = time.NewTimer(idleTimeout) + idle = timer.C + defer timer.Stop() + } + completed := 0 +wait: + for { + select { + case <-finished: + completed++ + break wait + case <-activity: + if timer != nil { + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(idleTimeout) + } + case <-idle: + break wait + } + } + + // Closing both packet endpoints interrupts their blocking reads. A read + // deadline interrupts the control watcher without removing the connection + // from lifecycle tracking before all association goroutines have exited. + _ = local.Close() + _ = upstream.Close() + _ = control.SetReadDeadline(time.Now()) + for completed < 3 { + <-finished + completed++ + } +} + func (s *SOCKS5Server) negotiate(conn net.Conn) error { methods, err := readSOCKSGreeting(conn) if err != nil { @@ -220,8 +594,11 @@ func readUserPasswordRequest(r io.Reader) (string, string, error) { } type socksRequest struct { - command byte - address string + command byte + addressType byte + host string + port uint16 + address string } type socksProtocolError struct { @@ -259,15 +636,18 @@ func readSOCKSRequest(r io.Reader) (socksRequest, error) { return socksRequest{}, err } port := binary.BigEndian.Uint16(portBytes[:]) - if port == 0 { + if port == 0 && header[1] != socksCommandUDP { return socksRequest{}, &socksProtocolError{ reply: socksReplyAddressUnsupported, err: errors.New("socks5: zero destination port"), } } return socksRequest{ - command: header[1], - address: net.JoinHostPort(host, fmt.Sprintf("%d", port)), + command: header[1], + addressType: header[3], + host: host, + port: port, + address: net.JoinHostPort(host, fmt.Sprintf("%d", port)), }, nil } @@ -317,6 +697,81 @@ func readSOCKSHost(r io.Reader, addressType byte) (string, error) { } } +func parseSOCKSUDPDatagram(packet []byte) ([]byte, string, error) { + if len(packet) > maxUDPDatagramSize { + return nil, "", errors.New("socks5: UDP datagram exceeds maximum size") + } + if len(packet) < 4 { + return nil, "", io.ErrUnexpectedEOF + } + if packet[0] != 0 || packet[1] != 0 { + return nil, "", errors.New("socks5: nonzero UDP reserved field") + } + if packet[2] != 0 { + return nil, "", errSOCKSUDPFragmented + } + + reader := bytes.NewReader(packet[4:]) + host, err := readSOCKSHost(reader, packet[3]) + if err != nil { + return nil, "", err + } + portBytes := [2]byte{} + if _, err := io.ReadFull(reader, portBytes[:]); err != nil { + return nil, "", err + } + port := binary.BigEndian.Uint16(portBytes[:]) + if port == 0 { + return nil, "", errors.New("socks5: zero UDP destination port") + } + payloadOffset := len(packet) - reader.Len() + return packet[payloadOffset:], net.JoinHostPort(host, strconv.Itoa(int(port))), nil +} + +func buildSOCKSUDPDatagram(payload []byte, address string) ([]byte, error) { + host, portString, err := net.SplitHostPort(address) + if err != nil { + return nil, err + } + port, err := strconv.ParseUint(portString, 10, 16) + if err != nil || port == 0 { + return nil, errors.New("socks5: invalid UDP source port") + } + + packet := make([]byte, 4, 4+net.IPv6len+2+len(payload)) + if ip := net.ParseIP(host); ip != nil { + if ip4 := ip.To4(); ip4 != nil { + packet[3] = socksAddressIPv4 + packet = append(packet, ip4...) + } else { + ip16 := ip.To16() + if ip16 == nil { + return nil, errors.New("socks5: invalid UDP source IP") + } + packet[3] = socksAddressIPv6 + packet = append(packet, ip16...) + } + } else { + if len(host) == 0 || len(host) > 255 { + return nil, errors.New("socks5: invalid UDP source host length") + } + for _, value := range []byte(host) { + if value <= 0x20 || value == 0x7f { + return nil, errors.New("socks5: invalid UDP source host") + } + } + packet[3] = socksAddressDomain + packet = append(packet, byte(len(host))) + packet = append(packet, host...) + } + packet = binary.BigEndian.AppendUint16(packet, uint16(port)) + if len(packet)+len(payload) > maxUDPDatagramSize { + return nil, errors.New("socks5: UDP datagram exceeds maximum size") + } + packet = append(packet, payload...) + return packet, nil +} + func writeSOCKSReply(w io.Writer, reply byte, address net.Addr) error { ip := net.IPv4zero port := 0 diff --git a/internal/proxy/socks5_udp_test.go b/internal/proxy/socks5_udp_test.go new file mode 100644 index 0000000..028ad4f --- /dev/null +++ b/internal/proxy/socks5_udp_test.go @@ -0,0 +1,642 @@ +package proxy + +import ( + "bufio" + "bytes" + "context" + "encoding/binary" + "errors" + "io" + "net" + "sync" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +func TestSOCKS5UDPAssociateRoundTrip(t *testing.T) { + echoAddress, stopEcho := startUDPEcho(t) + defer stopEcho() + + dialer := testPacketDialer{ + Dialer: directDialer(), + dialPacket: func(context.Context) (transport.PacketConn, error) { + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + return nil, err + } + return &directPacketConn{UDPConn: conn}, nil + }, + } + server, proxyAddress, stopProxy := startSOCKS5(t, Config{Dialer: dialer}) + defer stopProxy(server) + + control := dialTCP(t, proxyAddress) + defer control.Close() + socksGreeting(t, control, nil) + mustWrite(t, control, ipv4SOCKSRequest(socksCommandUDP, net.IPv4zero, 0)) + reply, relayAddress := readSOCKSReplyAddress(t, control) + if reply != socksReplySucceeded { + t.Fatalf("reply = %d, want success", reply) + } + if relayAddress.Port == 0 || !relayAddress.IP.Equal(net.IPv4(127, 0, 0, 1)) { + t.Fatalf("UDP relay address = %v, want loopback with a nonzero port", relayAddress) + } + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + _ = client.SetDeadline(time.Now().Add(3 * time.Second)) + payload := []byte("autocar UDP associate") + packet, err := buildSOCKSUDPDatagram(payload, echoAddress) + if err != nil { + t.Fatal(err) + } + if _, err := client.WriteToUDP(packet, relayAddress); err != nil { + t.Fatal(err) + } + + buffer := make([]byte, maxUDPDatagramSize) + n, source, err := client.ReadFromUDP(buffer) + if err != nil { + t.Fatal(err) + } + if !source.IP.Equal(relayAddress.IP) || source.Port != relayAddress.Port { + t.Fatalf("response source = %v, want %v", source, relayAddress) + } + got, gotSource, err := parseSOCKSUDPDatagram(buffer[:n]) + if err != nil { + t.Fatal(err) + } + if gotSource != echoAddress { + t.Fatalf("encapsulated source = %q, want %q", gotSource, echoAddress) + } + if !bytes.Equal(got, payload) { + t.Fatalf("payload = %q, want %q", got, payload) + } +} + +func TestSOCKS5UDPAssociateDropsFragmentsAndWrongSourcePort(t *testing.T) { + packetConn := newRecordingPacketConn() + dialer := testPacketDialer{ + Dialer: directDialer(), + dialPacket: func(context.Context) (transport.PacketConn, error) { + return packetConn, nil + }, + } + server, proxyAddress, stopProxy := startSOCKS5(t, Config{Dialer: dialer}) + defer stopProxy(server) + + allowed, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer allowed.Close() + wrongPort, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer wrongPort.Close() + + control := dialTCP(t, proxyAddress) + defer control.Close() + socksGreeting(t, control, nil) + mustWrite(t, control, ipv4SOCKSRequest( + socksCommandUDP, + net.IPv4zero, + uint16(allowed.LocalAddr().(*net.UDPAddr).Port), + )) + reply, relayAddress := readSOCKSReplyAddress(t, control) + if reply != socksReplySucceeded { + t.Fatalf("reply = %d, want success", reply) + } + + valid, err := buildSOCKSUDPDatagram([]byte("accepted"), "example.com:53") + if err != nil { + t.Fatal(err) + } + if _, err := wrongPort.WriteToUDP(valid, relayAddress); err != nil { + t.Fatal(err) + } + assertNoPacketSend(t, packetConn.sends) + + fragmented := append([]byte(nil), valid...) + fragmented[2] = 1 + if _, err := allowed.WriteToUDP(fragmented, relayAddress); err != nil { + t.Fatal(err) + } + assertNoPacketSend(t, packetConn.sends) + + if _, err := allowed.WriteToUDP(valid, relayAddress); err != nil { + t.Fatal(err) + } + select { + case sent := <-packetConn.sends: + if sent.address != "example.com:53" || string(sent.payload) != "accepted" { + t.Fatalf("upstream send = %#v", sent) + } + case <-time.After(time.Second): + t.Fatal("valid datagram was not sent upstream") + } + + packetConn.incoming <- packetRecord{payload: []byte("response"), address: "192.0.2.9:5353"} + _ = allowed.SetReadDeadline(time.Now().Add(time.Second)) + buffer := make([]byte, 128) + n, _, err := allowed.ReadFromUDP(buffer) + if err != nil { + t.Fatal(err) + } + payload, source, err := parseSOCKSUDPDatagram(buffer[:n]) + if err != nil { + t.Fatal(err) + } + if string(payload) != "response" || source != "192.0.2.9:5353" { + t.Fatalf("downstream packet = %q from %q", payload, source) + } +} + +func TestSOCKS5UDPAssociateHonorsTransportPayloadLimit(t *testing.T) { + recorded := newRecordingPacketConn() + upstream := &limitedPacketConn{recordingPacketConn: recorded, limit: 4} + dialer := testPacketDialer{ + Dialer: directDialer(), + dialPacket: func(context.Context) (transport.PacketConn, error) { + return upstream, nil + }, + } + server, proxyAddress, stopProxy := startSOCKS5(t, Config{Dialer: dialer}) + defer stopProxy(server) + + control := dialTCP(t, proxyAddress) + defer control.Close() + socksGreeting(t, control, nil) + mustWrite(t, control, ipv4SOCKSRequest(socksCommandUDP, net.IPv4zero, 0)) + reply, relayAddress := readSOCKSReplyAddress(t, control) + if reply != socksReplySucceeded { + t.Fatalf("reply = %d, want success", reply) + } + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + oversized, err := buildSOCKSUDPDatagram([]byte("12345"), "example.com:53") + if err != nil { + t.Fatal(err) + } + if _, err := client.WriteToUDP(oversized, relayAddress); err != nil { + t.Fatal(err) + } + assertNoPacketSend(t, recorded.sends) + + valid, err := buildSOCKSUDPDatagram([]byte("1234"), "example.com:53") + if err != nil { + t.Fatal(err) + } + if _, err := client.WriteToUDP(valid, relayAddress); err != nil { + t.Fatal(err) + } + select { + case sent := <-recorded.sends: + if string(sent.payload) != "1234" || sent.address != "example.com:53" { + t.Fatalf("upstream send = %#v", sent) + } + case <-time.After(time.Second): + t.Fatal("valid datagram was not sent after oversized datagram") + } +} + +func TestSOCKS5UDPAssociateRejectsMismatchedRequestedAddressBeforeUpstream(t *testing.T) { + dialCalled := make(chan struct{}, 1) + dialer := testPacketDialer{ + Dialer: directDialer(), + dialPacket: func(context.Context) (transport.PacketConn, error) { + dialCalled <- struct{}{} + return newRecordingPacketConn(), nil + }, + } + server, proxyAddress, stopProxy := startSOCKS5(t, Config{Dialer: dialer}) + defer stopProxy(server) + + control := dialTCP(t, proxyAddress) + defer control.Close() + socksGreeting(t, control, nil) + mustWrite(t, control, ipv4SOCKSRequest(socksCommandUDP, net.ParseIP("192.0.2.44"), 5353)) + reply, _ := readSOCKSReplyAddress(t, control) + if reply != socksReplyNotAllowed { + t.Fatalf("reply = %d, want not allowed", reply) + } + select { + case <-dialCalled: + t.Fatal("invalid UDP ASSOCIATE request allocated an upstream session") + default: + } +} + +func TestSOCKS5UDPAssociateShutdownClosesPacketConn(t *testing.T) { + packetConn := newRecordingPacketConn() + dialer := testPacketDialer{ + Dialer: directDialer(), + dialPacket: func(context.Context) (transport.PacketConn, error) { + return packetConn, nil + }, + } + server, proxyAddress, stopProxy := startSOCKS5(t, Config{Dialer: dialer}) + defer stopProxy(server) + control := dialTCP(t, proxyAddress) + defer control.Close() + socksGreeting(t, control, nil) + mustWrite(t, control, ipv4SOCKSRequest(socksCommandUDP, net.IPv4zero, 0)) + if reply, _ := readSOCKSReplyAddress(t, control); reply != socksReplySucceeded { + t.Fatalf("reply = %d, want success", reply) + } + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + if err := server.Shutdown(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("Shutdown error = %v, want deadline exceeded", err) + } + select { + case <-packetConn.closed: + case <-time.After(time.Second): + t.Fatal("packet connection was not closed by forced shutdown") + } + _ = control.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := control.Read(make([]byte, 1)); err == nil { + t.Fatal("control connection remained open after forced shutdown") + } +} + +func TestSOCKS5UDPAssociateIdleTimeout(t *testing.T) { + packetConn := newRecordingPacketConn() + dialer := testPacketDialer{ + Dialer: directDialer(), + dialPacket: func(context.Context) (transport.PacketConn, error) { + return packetConn, nil + }, + } + server, proxyAddress, stopProxy := startSOCKS5(t, Config{ + Dialer: dialer, + IdleTimeout: 30 * time.Millisecond, + }) + defer stopProxy(server) + control := dialTCP(t, proxyAddress) + defer control.Close() + socksGreeting(t, control, nil) + mustWrite(t, control, ipv4SOCKSRequest(socksCommandUDP, net.IPv4zero, 0)) + if reply, _ := readSOCKSReplyAddress(t, control); reply != socksReplySucceeded { + t.Fatalf("reply = %d, want success", reply) + } + select { + case <-packetConn.closed: + case <-time.After(time.Second): + t.Fatal("idle UDP association did not close") + } + _ = control.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := control.Read(make([]byte, 1)); err == nil { + t.Fatal("idle UDP control connection remained open") + } +} + +func TestSOCKSUDPClientEndpointSourceRestrictions(t *testing.T) { + endpoint := &socksUDPClientEndpoint{peerIP: net.ParseIP("192.0.2.1")} + if endpoint.accept(&net.UDPAddr{IP: net.ParseIP("192.0.2.2"), Port: 1000}) { + t.Fatal("accepted a datagram from a different IP") + } + if !endpoint.accept(&net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 1000}) { + t.Fatal("rejected first valid source") + } + if endpoint.accept(&net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 1001}) { + t.Fatal("accepted a different source port after locking") + } +} + +func TestValidateUDPAssociateRequestAddressAndPortSemantics(t *testing.T) { + domainLookup := func(_ context.Context, host string) ([]net.IPAddr, error) { + switch host { + case "client.example": + return []net.IPAddr{ + {IP: net.ParseIP("2001:db8::99")}, + {IP: net.ParseIP("192.0.2.10")}, + }, nil + case "other.example": + return []net.IPAddr{{IP: net.ParseIP("192.0.2.11")}}, nil + default: + return nil, &net.DNSError{Name: host, Err: "test lookup failure"} + } + } + tests := []struct { + name string + request []byte + peerIP net.IP + wantPort int + wantReply byte + }{ + { + name: "IPv4 concrete address and port", + request: ipv4SOCKSRequest(socksCommandUDP, net.ParseIP("192.0.2.10"), 5300), + peerIP: net.ParseIP("192.0.2.10"), + wantPort: 5300, + }, + { + name: "IPv4 unspecified dynamic port", + request: ipv4SOCKSRequest(socksCommandUDP, net.IPv4zero, 0), + peerIP: net.ParseIP("192.0.2.10"), + wantPort: 0, + }, + { + name: "IPv4 address mismatch", + request: ipv4SOCKSRequest(socksCommandUDP, net.ParseIP("192.0.2.11"), 5300), + peerIP: net.ParseIP("192.0.2.10"), + wantReply: socksReplyNotAllowed, + }, + { + name: "IPv6 concrete address and port", + request: ipv6SOCKSRequest(socksCommandUDP, net.ParseIP("2001:db8::10"), 5353), + peerIP: net.ParseIP("2001:db8::10"), + wantPort: 5353, + }, + { + name: "IPv6 unspecified dynamic port", + request: ipv6SOCKSRequest(socksCommandUDP, net.IPv6zero, 0), + peerIP: net.ParseIP("2001:db8::10"), + wantPort: 0, + }, + { + name: "IPv6 address mismatch", + request: ipv6SOCKSRequest(socksCommandUDP, net.ParseIP("2001:db8::11"), 5353), + peerIP: net.ParseIP("2001:db8::10"), + wantReply: socksReplyNotAllowed, + }, + { + name: "domain resolves to control peer", + request: domainSOCKSRequest(socksCommandUDP, "client.example", 6000), + peerIP: net.ParseIP("192.0.2.10"), + wantPort: 6000, + }, + { + name: "domain resolves elsewhere", + request: domainSOCKSRequest(socksCommandUDP, "other.example", 6000), + peerIP: net.ParseIP("192.0.2.10"), + wantReply: socksReplyNotAllowed, + }, + { + name: "domain lookup failure", + request: domainSOCKSRequest(socksCommandUDP, "missing.example", 6000), + peerIP: net.ParseIP("192.0.2.10"), + wantReply: socksReplyHostUnreachable, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + request, err := readSOCKSRequest(bytes.NewReader(test.request)) + if err != nil { + t.Fatal(err) + } + port, err := validateUDPAssociateRequest(context.Background(), request, test.peerIP, domainLookup) + if test.wantReply == 0 { + if err != nil || port != test.wantPort { + t.Fatalf("port=%d error=%v, want port %d", port, err, test.wantPort) + } + return + } + var protocolErr *socksProtocolError + if !errors.As(err, &protocolErr) || protocolErr.reply != test.wantReply { + t.Fatalf("error = %v, want SOCKS reply %d", err, test.wantReply) + } + }) + } +} + +func TestSOCKSUDPDatagramAddressTypesAndLimits(t *testing.T) { + for _, address := range []string{"192.0.2.1:53", "[2001:db8::1]:443", "example.com:5353"} { + t.Run(address, func(t *testing.T) { + packet, err := buildSOCKSUDPDatagram([]byte("data"), address) + if err != nil { + t.Fatal(err) + } + payload, gotAddress, err := parseSOCKSUDPDatagram(packet) + if err != nil { + t.Fatal(err) + } + if string(payload) != "data" || gotAddress != address { + t.Fatalf("round trip = %q to %q", payload, gotAddress) + } + }) + } + + maximumPayload := make([]byte, maxUDPDatagramSize-10) // IPv4 header is 10 bytes. + if _, err := buildSOCKSUDPDatagram(maximumPayload, "192.0.2.1:53"); err != nil { + t.Fatalf("maximum-size datagram: %v", err) + } + if _, err := buildSOCKSUDPDatagram(append(maximumPayload, 0), "192.0.2.1:53"); err == nil { + t.Fatal("oversize datagram was accepted") + } +} + +func TestParseSOCKSUDPDatagramRejectsMalformedPackets(t *testing.T) { + valid, err := buildSOCKSUDPDatagram([]byte("x"), "192.0.2.1:53") + if err != nil { + t.Fatal(err) + } + fragmented := append([]byte(nil), valid...) + fragmented[2] = 1 + if _, _, err := parseSOCKSUDPDatagram(fragmented); !errors.Is(err, errSOCKSUDPFragmented) { + t.Fatalf("fragment error = %v", err) + } + for _, packet := range [][]byte{ + nil, + {0, 0, 0}, + {1, 0, 0, socksAddressIPv4, 127, 0, 0, 1, 0, 53}, + {0, 0, 0, 0xff, 0, 53}, + {0, 0, 0, socksAddressIPv4, 127, 0, 0, 1, 0, 0}, + } { + if _, _, err := parseSOCKSUDPDatagram(packet); err == nil { + t.Fatalf("malformed packet %v was accepted", packet) + } + } +} + +func TestReadSOCKSUDPRequestAllowsZeroPort(t *testing.T) { + request, err := readSOCKSRequest(bytes.NewReader(ipv4SOCKSRequest( + socksCommandUDP, + net.IPv4zero, + 0, + ))) + if err != nil { + t.Fatal(err) + } + if request.address != "0.0.0.0:0" { + t.Fatalf("address = %q", request.address) + } + if request.addressType != socksAddressIPv4 || request.host != "0.0.0.0" || request.port != 0 { + t.Fatalf("parsed UDP endpoint = type %d host %q port %d", request.addressType, request.host, request.port) + } + if _, err := readSOCKSRequest(bytes.NewReader(ipv4SOCKSRequest( + socksCommandConnect, + net.IPv4zero, + 0, + ))); err == nil { + t.Fatal("CONNECT with port zero was accepted") + } +} + +func ipv6SOCKSRequest(command byte, ip net.IP, port uint16) []byte { + request := []byte{socksVersion, command, 0, socksAddressIPv6} + request = append(request, ip.To16()...) + return binary.BigEndian.AppendUint16(request, port) +} + +func FuzzParseSOCKSUDPDatagram(f *testing.F) { + seed, _ := buildSOCKSUDPDatagram([]byte("payload"), "example.com:53") + f.Add(seed) + f.Add([]byte{0, 0, 0, socksAddressIPv4, 127, 0, 0, 1, 0, 53}) + f.Fuzz(func(t *testing.T, packet []byte) { + _, _, _ = parseSOCKSUDPDatagram(packet) + }) +} + +type testPacketDialer struct { + transport.Dialer + dialPacket func(context.Context) (transport.PacketConn, error) +} + +func (d testPacketDialer) DialPacket(ctx context.Context) (transport.PacketConn, error) { + return d.dialPacket(ctx) +} + +type directPacketConn struct { + *net.UDPConn +} + +func (c *directPacketConn) Send(payload []byte, address string) error { + target, err := net.ResolveUDPAddr("udp", address) + if err != nil { + return err + } + _, err = c.WriteToUDP(payload, target) + return err +} + +func (c *directPacketConn) Receive() ([]byte, string, error) { + buffer := make([]byte, maxUDPDatagramSize) + n, source, err := c.ReadFromUDP(buffer) + if err != nil { + return nil, "", err + } + return buffer[:n], source.String(), nil +} + +type packetRecord struct { + payload []byte + address string +} + +type recordingPacketConn struct { + sends chan packetRecord + incoming chan packetRecord + closed chan struct{} + once sync.Once +} + +type limitedPacketConn struct { + *recordingPacketConn + limit int +} + +func (c *limitedPacketConn) MaxPayloadSize() int { return c.limit } + +func newRecordingPacketConn() *recordingPacketConn { + return &recordingPacketConn{ + sends: make(chan packetRecord, 8), + incoming: make(chan packetRecord, 8), + closed: make(chan struct{}), + } +} + +func (c *recordingPacketConn) Send(payload []byte, address string) error { + record := packetRecord{payload: append([]byte(nil), payload...), address: address} + select { + case c.sends <- record: + return nil + case <-c.closed: + return net.ErrClosed + } +} + +func (c *recordingPacketConn) Receive() ([]byte, string, error) { + select { + case packet := <-c.incoming: + return append([]byte(nil), packet.payload...), packet.address, nil + case <-c.closed: + return nil, "", net.ErrClosed + } +} + +func (c *recordingPacketConn) Close() error { + c.once.Do(func() { close(c.closed) }) + return nil +} + +func startUDPEcho(t *testing.T) (string, func()) { + t.Helper() + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + go func() { + defer close(done) + buffer := make([]byte, maxUDPDatagramSize) + for { + n, source, err := conn.ReadFromUDP(buffer) + if err != nil { + return + } + _, _ = conn.WriteToUDP(buffer[:n], source) + } + }() + return conn.LocalAddr().String(), func() { + _ = conn.Close() + <-done + } +} + +func ipv4SOCKSRequest(command byte, ip net.IP, port uint16) []byte { + request := []byte{socksVersion, command, 0, socksAddressIPv4} + request = append(request, ip.To4()...) + return append(request, byte(port>>8), byte(port)) +} + +func readSOCKSReplyAddress(t *testing.T, reader io.Reader) (byte, *net.UDPAddr) { + t.Helper() + buffered, ok := reader.(*bufio.Reader) + if !ok { + buffered = bufio.NewReader(reader) + } + header := make([]byte, 4) + mustReadFull(t, buffered, header) + if header[0] != socksVersion || header[2] != 0 { + t.Fatalf("invalid SOCKS reply header %v", header) + } + host, err := readSOCKSHost(buffered, header[3]) + if err != nil { + t.Fatal(err) + } + portBytes := make([]byte, 2) + mustReadFull(t, buffered, portBytes) + port := int(portBytes[0])<<8 | int(portBytes[1]) + return header[1], &net.UDPAddr{IP: net.ParseIP(host), Port: port} +} + +func assertNoPacketSend(t *testing.T, sends <-chan packetRecord) { + t.Helper() + select { + case packet := <-sends: + t.Fatalf("unexpected upstream packet %#v", packet) + case <-time.After(50 * time.Millisecond): + } +} diff --git a/internal/security/dialer.go b/internal/security/dialer.go index cc6a4a0..342faae 100644 --- a/internal/security/dialer.go +++ b/internal/security/dialer.go @@ -222,6 +222,68 @@ func (d *SafeDialer) DialContext(ctx context.Context, network, address string) ( return nil, fmt.Errorf("security: %s has no addresses matching %s", host, network) } +// ResolveUDPContext resolves a UDP destination and returns only numeric +// addresses that pass the same port, special-use, private-network and custom +// CIDR policy enforced by DialContext. Callers must use one of the returned +// addresses directly and must not resolve the original hostname again. This +// is used by datagram transports where a connected net.Conn is not suitable. +func (d *SafeDialer) ResolveUDPContext(ctx context.Context, address string) ([]netip.AddrPort, error) { + if d == nil { + return nil, errors.New("security: nil SafeDialer") + } + if ctx == nil { + return nil, errors.New("security: nil resolve context") + } + if err := context.Cause(ctx); err != nil { + return nil, err + } + host, portText, err := net.SplitHostPort(address) + if err != nil { + return nil, fmt.Errorf("security: invalid destination %q: %w", address, err) + } + if host == "" { + return nil, errors.New("security: destination host is required") + } + portNumber, err := strconv.ParseUint(portText, 10, 16) + if err != nil || portNumber == 0 { + return nil, fmt.Errorf("security: destination port %q is not a number from 1 to 65535", portText) + } + port := uint16(portNumber) + deniedPorts := d.deniedPorts + if deniedPorts == nil { + deniedPorts = map[uint16]struct{}{25: {}, 465: {}, 587: {}} + } + if _, denied := deniedPorts[port]; denied { + return nil, fmt.Errorf("%w: %d", ErrDeniedPort, port) + } + + addresses, err := d.resolve(ctx, "ip", host) + if err != nil { + return nil, err + } + unsafeCount := 0 + approved := make([]netip.AddrPort, 0, len(addresses)) + for _, addr := range addresses { + addr = addr.Unmap() + if err := validateDestinationIP(addr, d.allowPrivate); err != nil { + unsafeCount++ + continue + } + if matchesDeniedPrefix(addr, d.deniedNets) { + unsafeCount++ + continue + } + approved = append(approved, netip.AddrPortFrom(addr, port)) + } + if len(approved) != 0 { + return approved, nil + } + if unsafeCount != 0 { + return nil, fmt.Errorf("%w: %s resolved only to prohibited addresses", ErrUnsafeAddress, host) + } + return nil, fmt.Errorf("security: %s has no usable UDP addresses", host) +} + func (d *SafeDialer) resolve(ctx context.Context, network, host string) ([]netip.Addr, error) { if literal, err := netip.ParseAddr(host); err == nil { if literal.Zone() != "" { diff --git a/internal/security/dialer_test.go b/internal/security/dialer_test.go index 85cbfd9..57b1a30 100644 --- a/internal/security/dialer_test.go +++ b/internal/security/dialer_test.go @@ -23,6 +23,23 @@ type fakeResolver struct { calls []resolverCall } +type rotatingResolver struct { + mu sync.Mutex + answers [][]netip.Addr + calls []resolverCall +} + +func (r *rotatingResolver) LookupNetIP(_ context.Context, network, host string) ([]netip.Addr, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.calls = append(r.calls, resolverCall{network: network, host: host}) + index := len(r.calls) - 1 + if index >= len(r.answers) { + index = len(r.answers) - 1 + } + return append([]netip.Addr(nil), r.answers[index]...), nil +} + func (r *fakeResolver) LookupNetIP(_ context.Context, network, host string) ([]netip.Addr, error) { r.calls = append(r.calls, resolverCall{network: network, host: host}) return append([]netip.Addr(nil), r.addresses...), r.err @@ -408,6 +425,89 @@ func TestSafeDialerChecksCanceledContextBeforeDial(t *testing.T) { } } +func TestSafeDialerResolveUDPReturnsOnlyApprovedNumericAddresses(t *testing.T) { + resolver := &fakeResolver{addresses: []netip.Addr{ + netip.MustParseAddr("10.0.0.1"), + netip.MustParseAddr("8.8.8.8"), + netip.MustParseAddr("2606:4700:4700::1111"), + netip.MustParseAddr("8.8.8.8"), + }} + safe := NewSafeDialer(SafeDialerOptions{Resolver: resolver}) + addresses, err := safe.ResolveUDPContext(context.Background(), "dns.example:53") + if err != nil { + t.Fatal(err) + } + want := []netip.AddrPort{ + netip.MustParseAddrPort("8.8.8.8:53"), + netip.MustParseAddrPort("[2606:4700:4700::1111]:53"), + } + if !reflect.DeepEqual(addresses, want) { + t.Fatalf("UDP addresses = %v, want %v", addresses, want) + } + if !reflect.DeepEqual(resolver.calls, []resolverCall{{network: "ip", host: "dns.example"}}) { + t.Fatalf("resolver calls = %v", resolver.calls) + } +} + +func TestSafeDialerResolveUDPRejectsDeniedPortBeforeDNS(t *testing.T) { + resolver := &fakeResolver{addresses: []netip.Addr{netip.MustParseAddr("8.8.8.8")}} + safe := NewSafeDialer(SafeDialerOptions{Resolver: resolver}) + if _, err := safe.ResolveUDPContext(context.Background(), "mail.example:465"); !errors.Is(err, ErrDeniedPort) { + t.Fatalf("denied UDP port error = %v", err) + } + if len(resolver.calls) != 0 { + t.Fatalf("denied UDP port caused DNS resolution: %v", resolver.calls) + } +} + +func TestSafeDialerResolveUDPRejectsUnsafeResolution(t *testing.T) { + resolver := &fakeResolver{addresses: []netip.Addr{ + netip.MustParseAddr("127.0.0.1"), + netip.MustParseAddr("169.254.169.254"), + }} + safe := NewSafeDialer(SafeDialerOptions{Resolver: resolver, AllowPrivate: true}) + if _, err := safe.ResolveUDPContext(context.Background(), "internal.example:53"); !errors.Is(err, ErrUnsafeAddress) { + t.Fatalf("unsafe UDP resolution error = %v", err) + } +} + +func TestSafeDialerResolveUDPNumericLiteralSkipsDNS(t *testing.T) { + resolver := &fakeResolver{err: errors.New("must not resolve")} + safe := NewSafeDialer(SafeDialerOptions{Resolver: resolver}) + addresses, err := safe.ResolveUDPContext(context.Background(), "[2606:4700:4700::1111]:443") + if err != nil { + t.Fatal(err) + } + want := []netip.AddrPort{netip.MustParseAddrPort("[2606:4700:4700::1111]:443")} + if !reflect.DeepEqual(addresses, want) { + t.Fatalf("literal UDP addresses = %v, want %v", addresses, want) + } + if len(resolver.calls) != 0 { + t.Fatalf("numeric UDP literal caused DNS resolution: %v", resolver.calls) + } +} + +func TestSafeDialerResolveUDPFreezesDNSAnswerAgainstRebinding(t *testing.T) { + resolver := &rotatingResolver{answers: [][]netip.Addr{ + {netip.MustParseAddr("8.8.8.8")}, + {netip.MustParseAddr("127.0.0.1")}, + }} + safe := NewSafeDialer(SafeDialerOptions{Resolver: resolver}) + addresses, err := safe.ResolveUDPContext(context.Background(), "rebinding.example:53") + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(addresses, []netip.AddrPort{netip.MustParseAddrPort("8.8.8.8:53")}) { + t.Fatalf("frozen numeric answer = %v", addresses) + } + resolver.mu.Lock() + calls := append([]resolverCall(nil), resolver.calls...) + resolver.mu.Unlock() + if len(calls) != 1 { + t.Fatalf("ResolveUDPContext performed %d lookups, want exactly 1", len(calls)) + } +} + func formatPort(port uint16) string { const digits = "0123456789" if port == 0 { diff --git a/internal/transport/transport.go b/internal/transport/transport.go index 139c38f..4f5f5a5 100644 --- a/internal/transport/transport.go +++ b/internal/transport/transport.go @@ -13,6 +13,30 @@ type Dialer interface { DialContext(ctx context.Context, network, address string) (net.Conn, error) } +// PacketDialer is an optional capability implemented by transports that can +// carry datagrams. Proxy frontends must continue to work with a plain Dialer; +// datagram commands are advertised only when this interface is available. +type PacketDialer interface { + DialPacket(ctx context.Context) (PacketConn, error) +} + +// PacketConn carries independent datagrams through an authenticated tunnel. +// Send consumes payload before returning. Close must unblock a concurrent +// Receive call so proxy shutdown cannot leak goroutines. +type PacketConn interface { + Send(payload []byte, address string) error + Receive() (payload []byte, address string, err error) + Close() error +} + +// PacketPayloadSizer is an optional capability for packet transports with a +// logical-message limit below the UDP protocol maximum. Frontends use it to +// reject an oversized payload without tearing down an otherwise healthy +// association. +type PacketPayloadSizer interface { + MaxPayloadSize() int +} + // DialFunc adapts a function to Dialer. type DialFunc func(context.Context, string, string) (net.Conn, error) diff --git a/internal/tunnel/common.go b/internal/tunnel/common.go index d849689..4c610c1 100644 --- a/internal/tunnel/common.go +++ b/internal/tunnel/common.go @@ -21,10 +21,11 @@ import ( ) const ( - defaultHandshakeTimeout = 10 * time.Second - defaultDialTimeout = 10 * time.Second - defaultMaxStreams = 1024 - defaultMaxConnections = 256 + defaultHandshakeTimeout = 10 * time.Second + defaultDialTimeout = 10 * time.Second + defaultMaxStreams = 1024 + defaultMaxConnections = 256 + defaultMaxClientConnections = 32 ) // RemoteError is returned when the authenticated exit rejects a CONNECT diff --git a/internal/tunnel/tls.go b/internal/tunnel/tls.go index 27eb937..f5e05bc 100644 --- a/internal/tunnel/tls.go +++ b/internal/tunnel/tls.go @@ -7,6 +7,7 @@ import ( "fmt" "io" "net" + "net/netip" "sync" "time" @@ -23,6 +24,7 @@ type TLSServerConfig struct { HandshakeTimeout time.Duration DialTimeout time.Duration MaxConcurrentStreams int + MaxClientConnections int } // TLSServer serves one tunneled TCP stream per TLS 1.3 connection. It is a @@ -31,6 +33,7 @@ type TLSServer struct { listener net.Listener tlsConfig *tls.Config core *serverCore + clients *sourceConnectionLimiter ctx context.Context cancel context.CancelFunc @@ -59,6 +62,16 @@ func ListenTLS(config TLSServerConfig) (*TLSServer, error) { if err != nil { return nil, err } + maxClientConnections := config.MaxClientConnections + if maxClientConnections < 0 { + return nil, errors.New("tunnel: maximum TLS client connections cannot be negative") + } + if maxClientConnections == 0 { + maxClientConnections = min(defaultMaxClientConnections, cap(core.sem)) + } + if maxClientConnections > cap(core.sem) { + return nil, fmt.Errorf("tunnel: maximum TLS client connections (%d) exceeds maximum concurrent streams (%d)", maxClientConnections, cap(core.sem)) + } listener, err := net.Listen("tcp", config.Address) if err != nil { return nil, fmt.Errorf("tunnel: listen TLS fallback: %w", err) @@ -68,6 +81,7 @@ func ListenTLS(config TLSServerConfig) (*TLSServer, error) { listener: listener, tlsConfig: tlsConfig, core: core, + clients: newSourceConnectionLimiter(maxClientConnections), ctx: ctx, cancel: cancel, conns: make(map[net.Conn]struct{}), @@ -121,19 +135,27 @@ func (s *TLSServer) Serve(ctx context.Context) error { _ = raw.Close() continue } + sourceKey := tlsSourceKey(raw.RemoteAddr()) + if !s.clients.acquire(sourceKey) { + s.core.release() + s.lifecycle.Unlock() + _ = raw.Close() + continue + } tlsConn := tls.Server(raw, s.tlsConfig) s.connMu.Lock() s.conns[tlsConn] = struct{}{} s.connMu.Unlock() s.wg.Add(1) s.lifecycle.Unlock() - go s.serveTLSConnection(acceptCtx, tlsConn) + go s.serveTLSConnection(acceptCtx, tlsConn, sourceKey) } } -func (s *TLSServer) serveTLSConnection(ctx context.Context, conn *tls.Conn) { +func (s *TLSServer) serveTLSConnection(ctx context.Context, conn *tls.Conn, sourceKey string) { defer s.wg.Done() defer s.core.release() + defer s.clients.release(sourceKey) defer func() { s.connMu.Lock() delete(s.conns, conn) @@ -149,6 +171,73 @@ func (s *TLSServer) serveTLSConnection(ctx context.Context, conn *tls.Conn) { s.core.handleStream(ctx, conn, nil) } +type sourceConnectionLimiter struct { + mu sync.Mutex + limit int + active map[string]int +} + +func newSourceConnectionLimiter(limit int) *sourceConnectionLimiter { + return &sourceConnectionLimiter{limit: limit, active: make(map[string]int)} +} + +func (l *sourceConnectionLimiter) acquire(key string) bool { + l.mu.Lock() + defer l.mu.Unlock() + if l.active[key] >= l.limit { + return false + } + l.active[key]++ + return true +} + +func (l *sourceConnectionLimiter) release(key string) { + l.mu.Lock() + defer l.mu.Unlock() + if l.active[key] <= 1 { + delete(l.active, key) + return + } + l.active[key]-- +} + +func (l *sourceConnectionLimiter) count(key string) int { + l.mu.Lock() + defer l.mu.Unlock() + return l.active[key] +} + +func tlsSourceKey(address net.Addr) string { + if address == nil { + return "" + } + if tcpAddress, ok := address.(*net.TCPAddr); ok { + if ip, ok := netip.AddrFromSlice(tcpAddress.IP); ok { + return sourceIPKey(ip) + } + } + host, _, err := net.SplitHostPort(address.String()) + if err == nil { + if ip, parseErr := netip.ParseAddr(host); parseErr == nil { + return sourceIPKey(ip) + } + } + // Unknown address representations share one conservative bucket. Including + // an unparsed port here would let a peer obtain a fresh bucket per socket. + return "" +} + +func sourceIPKey(ip netip.Addr) string { + ip = ip.Unmap().WithZone("") + if ip.Is4() { + return ip.String() + } + if ip.Is6() { + return netip.PrefixFrom(ip, 64).Masked().String() + } + return "" +} + // Close stops the listener, closes active connections and waits for relays. func (s *TLSServer) Close() error { s.closeOnce.Do(func() { diff --git a/internal/tunnel/tls_limits_test.go b/internal/tunnel/tls_limits_test.go new file mode 100644 index 0000000..db6f6c0 --- /dev/null +++ b/internal/tunnel/tls_limits_test.go @@ -0,0 +1,233 @@ +package tunnel + +import ( + "context" + "net" + "testing" + "time" +) + +func TestTLSServerClientConnectionLimitValidationAndDefaults(t *testing.T) { + serverTLS, _ := testTLSConfigs(t) + + for _, test := range []struct { + name string + streams int + clients int + wantClient int + }{ + {name: "default below 32", streams: 2, wantClient: 2}, + {name: "default capped at 32", streams: 64, wantClient: 32}, + {name: "explicit", streams: 4, clients: 3, wantClient: 3}, + } { + t.Run(test.name, func(t *testing.T) { + server, err := ListenTLS(TLSServerConfig{ + Address: "127.0.0.1:0", + Token: testToken, + TLSConfig: serverTLS, + MaxConcurrentStreams: test.streams, + MaxClientConnections: test.clients, + }) + if err != nil { + t.Fatal(err) + } + defer server.Close() + if got := server.clients.limit; got != test.wantClient { + t.Fatalf("client connection limit = %d, want %d", got, test.wantClient) + } + }) + } + + for _, test := range []struct { + name string + streams int + clients int + }{ + {name: "negative", streams: 4, clients: -1}, + {name: "exceeds global limit", streams: 2, clients: 3}, + } { + t.Run(test.name, func(t *testing.T) { + server, err := ListenTLS(TLSServerConfig{ + Address: "127.0.0.1:0", + Token: testToken, + TLSConfig: serverTLS, + MaxConcurrentStreams: test.streams, + MaxClientConnections: test.clients, + }) + if server != nil { + _ = server.Close() + } + if err == nil { + t.Fatal("invalid client connection limit was accepted") + } + }) + } +} + +func TestTLSSourceConnectionLimiter(t *testing.T) { + limiter := newSourceConnectionLimiter(1) + first := tlsSourceKey(&net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 1000}) + same := tlsSourceKey(&net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 2000}) + other := tlsSourceKey(&net.TCPAddr{IP: net.ParseIP("192.0.2.11"), Port: 1000}) + if first != same { + t.Fatalf("IPv4 source key varies by port: %q != %q", first, same) + } + if first == other { + t.Fatalf("distinct IPv4 sources share key %q", first) + } + if !limiter.acquire(first) { + t.Fatal("first source acquisition failed") + } + if limiter.acquire(same) { + t.Fatal("same source exceeded its limit") + } + if !limiter.acquire(other) { + t.Fatal("different source was incorrectly limited") + } + if got := limiter.count(first); got != 1 { + t.Fatalf("first source count = %d, want 1", got) + } + limiter.release(first) + if got := limiter.count(first); got != 0 { + t.Fatalf("released source count = %d, want 0", got) + } + if !limiter.acquire(first) { + t.Fatal("released source slot was not reusable") + } +} + +func TestTLSSourceKeyUsesIPv6Prefix(t *testing.T) { + first := tlsSourceKey(&net.TCPAddr{IP: net.ParseIP("2001:db8:1:2::1"), Port: 1000}) + samePrefix := tlsSourceKey(&net.TCPAddr{IP: net.ParseIP("2001:db8:1:2::ffff"), Port: 2000}) + otherPrefix := tlsSourceKey(&net.TCPAddr{IP: net.ParseIP("2001:db8:1:3::1"), Port: 1000}) + if first != samePrefix { + t.Fatalf("IPv6 addresses in one /64 have different keys: %q != %q", first, samePrefix) + } + if first == otherPrefix { + t.Fatalf("distinct IPv6 /64 prefixes share key %q", first) + } +} + +func TestTLSSourceKeyFailsClosed(t *testing.T) { + for _, address := range []net.Addr{ + nil, + testTLSAddr("unparseable-one:1234"), + testTLSAddr("unparseable-two:5678"), + } { + if got := tlsSourceKey(address); got != "" { + t.Fatalf("unknown address %v received separate source key %q", address, got) + } + } +} + +type testTLSAddr string + +func (a testTLSAddr) Network() string { return "test" } +func (a testTLSAddr) String() string { return string(a) } + +func TestTLSServerClientLimitAcrossAcceptedConnections(t *testing.T) { + serverTLS, _ := testTLSConfigs(t) + server, err := ListenTLS(TLSServerConfig{ + Address: "127.0.0.1:0", + Token: testToken, + TLSConfig: serverTLS, + HandshakeTimeout: 30 * time.Second, + MaxConcurrentStreams: 2, + MaxClientConnections: 1, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + serveDone := make(chan error, 1) + go func() { serveDone <- server.Serve(ctx) }() + t.Cleanup(func() { + cancel() + _ = server.Close() + select { + case err := <-serveDone: + if err != nil { + t.Errorf("TLS Serve: %v", err) + } + case <-time.After(5 * time.Second): + t.Error("TLS Serve did not stop") + } + }) + + dialFrom := func(sourceIP string) net.Conn { + t.Helper() + dialer := net.Dialer{ + Timeout: 2 * time.Second, + LocalAddr: &net.TCPAddr{IP: net.ParseIP(sourceIP)}, + } + conn, err := dialer.Dial("tcp", server.Addr().String()) + if err != nil { + t.Fatalf("dial from %s: %v", sourceIP, err) + } + t.Cleanup(func() { _ = conn.Close() }) + return conn + } + + first := dialFrom("127.0.0.1") + firstKey := tlsSourceKey(first.LocalAddr()) + waitForTLSLimitState(t, "first connection admission", func() bool { + return server.clients.count(firstKey) == 1 && len(server.core.sem) == 1 + }) + + rejected := dialFrom("127.0.0.1") + if err := rejected.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatal(err) + } + buffer := make([]byte, 1) + if _, err := rejected.Read(buffer); err == nil { + t.Fatal("second connection from the same source remained open") + } else if netError, ok := err.(net.Error); ok && netError.Timeout() { + t.Fatal("second connection from the same source was not rejected") + } + waitForTLSLimitState(t, "same-source rejection rollback", func() bool { + return server.clients.count(firstKey) == 1 && len(server.core.sem) == 1 + }) + + different := dialFrom("127.0.0.2") + differentKey := tlsSourceKey(different.LocalAddr()) + if differentKey == firstKey { + t.Fatalf("different loopback sources share key %q", firstKey) + } + waitForTLSLimitState(t, "different-source admission", func() bool { + return server.clients.count(differentKey) == 1 && len(server.core.sem) == 2 + }) + + if err := first.Close(); err != nil { + t.Fatal(err) + } + waitForTLSLimitState(t, "first connection release", func() bool { + return server.clients.count(firstKey) == 0 && len(server.core.sem) == 1 + }) + + replacement := dialFrom("127.0.0.1") + waitForTLSLimitState(t, "released slot reuse", func() bool { + return server.clients.count(firstKey) == 1 && len(server.core.sem) == 2 + }) + + if err := replacement.Close(); err != nil { + t.Fatal(err) + } + if err := different.Close(); err != nil { + t.Fatal(err) + } + waitForTLSLimitState(t, "all connection releases", func() bool { + return server.clients.count(firstKey) == 0 && + server.clients.count(differentKey) == 0 && len(server.core.sem) == 0 + }) +} + +func waitForTLSLimitState(t *testing.T, description string, condition func() bool) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for !condition() { + if time.Now().After(deadline) { + t.Fatalf("timed out waiting for %s", description) + } + time.Sleep(5 * time.Millisecond) + } +} diff --git a/scripts/check-fork-provenance.sh b/scripts/check-fork-provenance.sh new file mode 100755 index 0000000..f0011a7 --- /dev/null +++ b/scripts/check-fork-provenance.sh @@ -0,0 +1,47 @@ +#!/usr/bin/env sh +set -eu + +expected_hysteria='v2.12.1' +expected_quic='v0.61.1-0.20260806010916-184d081eef3e' + +module_version() { + awk -v module="$2" ' + $1 == module { print $2; found = 1 } + END { if (!found) exit 1 } + ' "$1" +} + +replace_target() { + awk -v module="$2" ' + $1 == "replace" && $2 == module && $3 == "=>" { print $4; found = 1 } + END { if (!found) exit 1 } + ' "$1" +} + +root_hysteria="$(module_version go.mod github.com/apernet/hysteria/core/v2)" +root_quic="$(module_version go.mod github.com/apernet/quic-go)" +core_quic="$(module_version third_party/hysteria-core/go.mod github.com/apernet/quic-go)" +root_hysteria_replace="$(replace_target go.mod github.com/apernet/hysteria/core/v2)" +root_quic_replace="$(replace_target go.mod github.com/apernet/quic-go)" +core_quic_replace="$(replace_target third_party/hysteria-core/go.mod github.com/apernet/quic-go)" + +if [ "$root_hysteria" != "$expected_hysteria" ]; then + echo "local Hysteria fork provenance mismatch: go.mod requires $root_hysteria, source is $expected_hysteria" >&2 + echo "rebase and re-audit third_party/hysteria-core before changing the requirement" >&2 + exit 1 +fi +if [ "$root_quic" != "$expected_quic" ] || [ "$core_quic" != "$expected_quic" ]; then + echo "local QUIC fork provenance mismatch: root=$root_quic core=$core_quic source=$expected_quic" >&2 + echo "rebase and re-audit third_party/quic-go before changing either requirement" >&2 + exit 1 +fi +if [ "$root_hysteria_replace" != './third_party/hysteria-core' ] || \ + [ "$root_quic_replace" != './third_party/quic-go' ] || \ + [ "$core_quic_replace" != '../quic-go' ]; then + echo "local fork replacements are missing or redirected" >&2 + echo "root Hysteria=$root_hysteria_replace root QUIC=$root_quic_replace core QUIC=$core_quic_replace" >&2 + exit 1 +fi + +grep -Fq "core/v2\` $expected_hysteria" third_party/hysteria-core/AUTOCAR_PATCHES.md +grep -Fq "\`$expected_quic\`" third_party/quic-go/AUTOCAR_PATCHES.md diff --git a/scripts/govulncheck.sh b/scripts/govulncheck.sh new file mode 100755 index 0000000..9051567 --- /dev/null +++ b/scripts/govulncheck.sh @@ -0,0 +1,45 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +# GO-2026-5288 is an automatically generated Go DB entry that currently says +# all Hysteria v2 versions are affected. The reviewed upstream GHSA limits the +# vulnerable range to <= 2.8.1, while AutoCAR pins 2.12.1. The vulnerable +# feature is application-layer protocol sniffing, which this adapter neither +# imports nor enables. Keep the exception fail-closed: changing the version, +# introducing sniff hooks, finding another reachable advisory, receiving an +# incomplete JSON stream, or a scanner execution error all fail this script. +# +# https://github.com/advisories/GHSA-9fw6-xgg2-mq9q +# https://pkg.go.dev/vuln/GO-2026-5288 + +EXPECTED_HYSTERIA_VERSION=v2.12.1 +ACTUAL_HYSTERIA_VERSION=$(go list -m -f '{{.Version}}' github.com/apernet/hysteria/core/v2) +if [[ ${ACTUAL_HYSTERIA_VERSION} != "${EXPECTED_HYSTERIA_VERSION}" ]]; then + echo "error: review the GO-2026-5288 exception before changing Hysteria (${ACTUAL_HYSTERIA_VERSION})" >&2 + exit 1 +fi + +if find cmd internal -type f -name '*.go' \ + -exec grep -EHin 'RequestHook[[:space:]]*:|/sniff(["`])' {} +; then + echo "error: protocol sniffing is incompatible with the narrow GO-2026-5288 exception" >&2 + exit 1 +fi + +TOOL_DIR=${RUNNER_TEMP:-/tmp}/autocar-govulncheck-v1.7.0 +mkdir -p "${TOOL_DIR}" +GOBIN="${TOOL_DIR}" go install golang.org/x/vuln/cmd/govulncheck@v1.7.0 + +set +e +"${TOOL_DIR}/govulncheck" -format=json ./... | go run ./tools/vulnfilter +PIPELINE_STATUS=("${PIPESTATUS[@]}") +set -e + +SCANNER_STATUS=${PIPELINE_STATUS[0]} +FILTER_STATUS=${PIPELINE_STATUS[1]} +if (( FILTER_STATUS != 0 )); then + exit "${FILTER_STATUS}" +fi +if (( SCANNER_STATUS != 0 && SCANNER_STATUS != 3 )); then + echo "error: govulncheck failed with status ${SCANNER_STATUS}" >&2 + exit "${SCANNER_STATUS}" +fi diff --git a/scripts/netem-integration.sh b/scripts/netem-integration.sh index ca83260..e8848fc 100755 --- a/scripts/netem-integration.sh +++ b/scripts/netem-integration.sh @@ -39,6 +39,7 @@ SERVER_DEV="acs${RUN_ID}" CLIENT_IP=10.203.0.1 SERVER_IP=10.203.0.2 RELAY_PORT=7443 +RENO_RELAY_PORT=7445 UNUSED_QUIC_PORT=7444 BENCH_PORT=9000 ORIGIN_PORT=9080 @@ -55,6 +56,13 @@ SHORT_FLOW_BYTES=${AUTOCAR_SHORT_FLOW_BYTES:-131072} SHORT_FLOW_ITERATIONS=${AUTOCAR_SHORT_FLOW_ITERATIONS:-9} SHORT_FLOW_WARMUP=${AUTOCAR_SHORT_FLOW_WARMUP:-3} MIN_SHORT_FLOW_RATIO=${AUTOCAR_MIN_SHORT_FLOW_RATIO:-1.10} +MIN_BBR_RENO_RATIO=${AUTOCAR_MIN_BBR_RENO_RATIO:-1.10} +MIN_BRUTAL_TARGET_RATIO=${AUTOCAR_MIN_BRUTAL_TARGET_RATIO:-0.50} +BRUTAL_SERVER_UPLOAD_MBPS=15 +BRUTAL_SERVER_DOWNLOAD_MBPS=15 +BRUTAL_CLIENT_UPLOAD_MBPS=20 +BRUTAL_CLIENT_DOWNLOAD_MBPS=20 +BRUTAL_EXPECTED_TX_BYTES_SEC=1875000 PIDS=() @@ -174,11 +182,28 @@ start_background "${ARTIFACT_DIR}/relay.log" \ --cert="${WORK_DIR}/server.crt" \ --key="${WORK_DIR}/server.key" \ --token-file="${WORK_DIR}/relay-token" \ + --max-upload-mbps="${BRUTAL_SERVER_UPLOAD_MBPS}" \ + --max-download-mbps="${BRUTAL_SERVER_DOWNLOAD_MBPS}" \ + --allow-client-bandwidth \ --allow-private --deny-ports=none RELAY_PID=${STARTED_PID} -wait_for_log "${RELAY_PID}" "${ARTIFACT_DIR}/relay.log" "transport=quic" +wait_for_log "${RELAY_PID}" "${ARTIFACT_DIR}/relay.log" "transport=hy2" wait_for_log "${RELAY_PID}" "${ARTIFACT_DIR}/relay.log" "transport=tls" +# A second, otherwise identical relay makes the download comparison exercise +# the relay-side sender. Client flags alone only select the client-side sender. +start_background "${ARTIFACT_DIR}/relay-reno.log" \ + ip netns exec "${SERVER_NS}" "${AUTOCAR_BIN}" server \ + --listen="${SERVER_IP}:${RENO_RELAY_PORT}" \ + --tcp-listen="${SERVER_IP}:${RENO_RELAY_PORT}" \ + --cert="${WORK_DIR}/server.crt" \ + --key="${WORK_DIR}/server.key" \ + --token-file="${WORK_DIR}/relay-token" \ + --congestion=reno \ + --allow-private --deny-ports=none +RENO_RELAY_PID=${STARTED_PID} +wait_for_log "${RENO_RELAY_PID}" "${ARTIFACT_DIR}/relay-reno.log" "transport=hy2" + COMMON_BENCH=( --target="${SERVER_IP}:${BENCH_PORT}" --mode=download @@ -196,21 +221,65 @@ TUNNEL_AUTH=( --dial-timeout=3s --open-timeout=5s ) +TUNNEL_AUTH_RENO=( + --server="${SERVER_IP}:${RENO_RELAY_PORT}" + --server-name="${SERVER_IP}" + --ca="${WORK_DIR}/server.crt" + --token-file="${WORK_DIR}/relay-token" + --dial-timeout=3s + --open-timeout=5s +) run_client --transport=direct "${COMMON_BENCH[@]}" >"${ARTIFACT_DIR}/direct.json" run_client --transport=quic "${TUNNEL_AUTH[@]}" "${COMMON_BENCH[@]}" >"${ARTIFACT_DIR}/quic.json" run_client --transport=tls "${TUNNEL_AUTH[@]}" "${COMMON_BENCH[@]}" >"${ARTIFACT_DIR}/tls.json" -run_client --transport=auto \ - --server="${SERVER_IP}:${UNUSED_QUIC_PORT}" \ - --fallback-server="${SERVER_IP}:${RELAY_PORT}" \ - --server-name="${SERVER_IP}" \ - --ca="${WORK_DIR}/server.crt" \ - --token-file="${WORK_DIR}/relay-token" \ - --dial-timeout=1s --quic-attempt-timeout=1s --open-timeout=5s \ - --target="${SERVER_IP}:${BENCH_PORT}" --mode=download \ - --bytes=131072 --iterations=2 --warmup=0 --timeout=20s --json \ - >"${ARTIFACT_DIR}/auto-fallback.json" +# Prove both controller paths with real authenticated transfers while the +# original delay, random loss and rate limits are still active. The BBR run +# declares no bandwidth and therefore must negotiate a zero Tx rate. The +# Brutal run declares 20 Mbit/s in both directions, while the relay's 15 +# Mbit/s upload cap deterministically limits client-to-relay Tx to 1,875,000 +# bytes/s. Repeated uploads isolate the client-side sender so the congestion +# flag deterministically selects the controller under comparison. +MODE_PROOF_BENCH=( + --target="${SERVER_IP}:${BENCH_PORT}" + --mode=upload + --bytes=4194304 + --iterations=3 + --warmup=1 + --timeout=60s + --json +) +run_client --transport=quic --congestion=bbr --bbr-profile=standard \ + "${TUNNEL_AUTH[@]}" "${MODE_PROOF_BENCH[@]}" \ + >"${ARTIFACT_DIR}/bbr.json" +run_client --transport=quic --congestion=reno \ + "${TUNNEL_AUTH[@]}" "${MODE_PROOF_BENCH[@]}" \ + >"${ARTIFACT_DIR}/reno.json" +run_client --transport=quic --congestion=bbr --bbr-profile=standard \ + --upload-mbps="${BRUTAL_CLIENT_UPLOAD_MBPS}" \ + --download-mbps="${BRUTAL_CLIENT_DOWNLOAD_MBPS}" \ + "${TUNNEL_AUTH[@]}" "${MODE_PROOF_BENCH[@]}" \ + >"${ARTIFACT_DIR}/brutal.json" + +# Repeat the BBR/Reno proof in the opposite direction. Both clients use the +# same configuration; only the relay sender differs (default BBR vs the +# explicitly configured Reno relay above). +MODE_PROOF_DOWNLOAD=( + --target="${SERVER_IP}:${BENCH_PORT}" + --mode=download + --bytes=4194304 + --iterations=3 + --warmup=1 + --timeout=60s + --json +) +run_client --transport=quic --congestion=bbr --bbr-profile=standard \ + "${TUNNEL_AUTH[@]}" "${MODE_PROOF_DOWNLOAD[@]}" \ + >"${ARTIFACT_DIR}/bbr-download.json" +run_client --transport=quic --congestion=bbr --bbr-profile=standard \ + "${TUNNEL_AUTH_RENO[@]}" "${MODE_PROOF_DOWNLOAD[@]}" \ + >"${ARTIFACT_DIR}/reno-download.json" # This intentionally narrow acceleration profile isolates the benefit of a # warm, shared congestion-control context. Every direct iteration creates a @@ -222,6 +291,23 @@ ip netns exec "${CLIENT_NS}" tc qdisc replace dev "${CLIENT_DEV}" root netem \ ip netns exec "${SERVER_NS}" tc qdisc replace dev "${SERVER_DEV}" root netem \ delay "${DELAY_MS}ms" rate "${RATE}" limit 10000 +# Test cold QUIC-to-TLS fallback after the lossy throughput profile has +# completed. Keeping the latency and rate constraints while removing random +# loss isolates the fallback state machine from a coincidental dropped TLS +# handshake. The deadlines remain finite and the command must still succeed +# on its first attempt. +sleep 0.25 +run_client --transport=auto \ + --server="${SERVER_IP}:${UNUSED_QUIC_PORT}" \ + --fallback-server="${SERVER_IP}:${RELAY_PORT}" \ + --server-name="${SERVER_IP}" \ + --ca="${WORK_DIR}/server.crt" \ + --token-file="${WORK_DIR}/relay-token" \ + --dial-timeout=3s --quic-attempt-timeout=2s --open-timeout=8s \ + --target="${SERVER_IP}:${BENCH_PORT}" --mode=download \ + --bytes=131072 --iterations=2 --warmup=0 --timeout=30s --json \ + >"${ARTIFACT_DIR}/auto-fallback.json" + SHORT_BENCH=( --target="${SERVER_IP}:${BENCH_PORT}" --mode=download @@ -268,7 +354,7 @@ kill -0 "${ORIGIN_PID}" start_background "${ARTIFACT_DIR}/client-proxy.log" \ ip netns exec "${CLIENT_NS}" "${AUTOCAR_BIN}" client \ --transport=auto "${TUNNEL_AUTH[@]}" \ - --dial-timeout=1s --quic-attempt-timeout=1s --open-timeout=2s \ + --dial-timeout=2s --quic-attempt-timeout=2s --open-timeout=6s \ --socks= --http="127.0.0.1:${PROXY_PORT}" --https= PROXY_PID=${STARTED_PID} wait_for_log "${PROXY_PID}" "${ARTIFACT_DIR}/client-proxy.log" "local proxy started" @@ -349,26 +435,123 @@ fi python3 - "${ARTIFACT_DIR}/direct.json" "${ARTIFACT_DIR}/quic.json" \ "${ARTIFACT_DIR}/tls.json" "${ARTIFACT_DIR}/short-direct.json" \ - "${ARTIFACT_DIR}/short-quic.json" "${ARTIFACT_DIR}/summary.json" \ + "${ARTIFACT_DIR}/short-quic.json" "${ARTIFACT_DIR}/bbr.json" \ + "${ARTIFACT_DIR}/reno.json" "${ARTIFACT_DIR}/brutal.json" \ + "${ARTIFACT_DIR}/bbr-download.json" "${ARTIFACT_DIR}/reno-download.json" \ + "${ARTIFACT_DIR}/summary.json" \ "${DELAY_MS}" "${LOSS}" "${RATE}" "${SHORT_FLOW_BYTES}" \ - "${MIN_SHORT_FLOW_RATIO}" <<'PY' + "${MIN_SHORT_FLOW_RATIO}" "${MIN_BBR_RENO_RATIO}" \ + "${MIN_BRUTAL_TARGET_RATIO}" "${BRUTAL_SERVER_UPLOAD_MBPS}" \ + "${BRUTAL_SERVER_DOWNLOAD_MBPS}" "${BRUTAL_CLIENT_UPLOAD_MBPS}" \ + "${BRUTAL_CLIENT_DOWNLOAD_MBPS}" "${BRUTAL_EXPECTED_TX_BYTES_SEC}" <<'PY' import json import pathlib import sys -direct_path, quic_path, tls_path, short_direct_path, short_quic_path, output_path = map( - pathlib.Path, sys.argv[1:7] +( + direct_path, + quic_path, + tls_path, + short_direct_path, + short_quic_path, + bbr_path, + reno_path, + brutal_path, + bbr_download_path, + reno_download_path, + output_path, +) = map( + pathlib.Path, sys.argv[1:12] ) -delay_ms, loss, rate, short_flow_bytes, minimum_ratio = sys.argv[7:12] +( + delay_ms, + loss, + rate, + short_flow_bytes, + minimum_ratio, + minimum_bbr_reno_ratio, + minimum_brutal_target_ratio, + server_upload_mbps, + server_download_mbps, + client_upload_mbps, + client_download_mbps, + expected_brutal_tx, +) = sys.argv[12:24] direct = json.loads(direct_path.read_text()) quic = json.loads(quic_path.read_text()) tls = json.loads(tls_path.read_text()) short_direct = json.loads(short_direct_path.read_text()) short_quic = json.loads(short_quic_path.read_text()) +bbr = json.loads(bbr_path.read_text()) +reno = json.loads(reno_path.read_text()) +brutal = json.loads(brutal_path.read_text()) +bbr_download = json.loads(bbr_download_path.read_text()) +reno_download = json.loads(reno_download_path.read_text()) for name, result in (("direct", direct), ("quic", quic), ("tls", tls)): if result["median_mbps"] <= 0: raise SystemExit(f"{name} benchmark reported non-positive goodput") +for name, result in (("bbr", bbr), ("reno", reno), ("brutal", brutal)): + if result["median_mbps"] <= 0: + raise SystemExit(f"{name} controller proof reported non-positive goodput") +for name, result in (("bbr download", bbr_download), ("reno download", reno_download)): + if result["median_mbps"] <= 0: + raise SystemExit(f"{name} controller proof reported non-positive goodput") + +if bbr.get("acceleration") != "bbr-standard": + raise SystemExit( + f"BBR proof reported acceleration={bbr.get('acceleration')!r}, " + "want 'bbr-standard'" + ) +if bbr.get("negotiated_tx_bytes_per_second", 0) != 0: + raise SystemExit( + "BBR proof unexpectedly negotiated a non-zero Tx bandwidth: " + f"{bbr.get('negotiated_tx_bytes_per_second')!r}" + ) +if reno.get("acceleration") != "reno": + raise SystemExit( + f"Reno proof reported acceleration={reno.get('acceleration')!r}, " + "want 'reno'" + ) +if reno.get("negotiated_tx_bytes_per_second", 0) != 0: + raise SystemExit( + "Reno proof unexpectedly negotiated a non-zero Tx bandwidth: " + f"{reno.get('negotiated_tx_bytes_per_second')!r}" + ) + +bbr_reno_ratio = bbr["median_mbps"] / reno["median_mbps"] +if bbr_reno_ratio < float(minimum_bbr_reno_ratio): + raise SystemExit( + f"lossy BBR/Reno upload ratio {bbr_reno_ratio:.3f} is below " + f"the declared acceptance threshold {float(minimum_bbr_reno_ratio):.3f}" + ) + +bbr_reno_download_ratio = bbr_download["median_mbps"] / reno_download["median_mbps"] +if bbr_reno_download_ratio < float(minimum_bbr_reno_ratio): + raise SystemExit( + f"lossy BBR/Reno download ratio {bbr_reno_download_ratio:.3f} is below " + f"the declared acceptance threshold {float(minimum_bbr_reno_ratio):.3f}" + ) + +expected_brutal_tx = int(expected_brutal_tx) +if brutal.get("acceleration") != "brutal": + raise SystemExit( + f"Brutal proof reported acceleration={brutal.get('acceleration')!r}, " + "want 'brutal'" + ) +if brutal.get("negotiated_tx_bytes_per_second") != expected_brutal_tx: + raise SystemExit( + "Brutal proof negotiated Tx bandwidth " + f"{brutal.get('negotiated_tx_bytes_per_second')!r}, " + f"want {expected_brutal_tx} bytes/s" + ) +brutal_target_mbps = expected_brutal_tx * 8 / 1_000_000 +brutal_target_ratio = brutal["median_mbps"] / brutal_target_mbps +if brutal_target_ratio < float(minimum_brutal_target_ratio): + raise SystemExit( + f"Brutal achieved/target ratio {brutal_target_ratio:.3f} is below " + f"the declared acceptance threshold {float(minimum_brutal_target_ratio):.3f}" + ) short_ratio = short_quic["median_mbps"] / short_direct["median_mbps"] if short_ratio < float(minimum_ratio): @@ -392,6 +575,54 @@ summary = { "quic": quic["median_mbps"] / direct["median_mbps"], "tls": tls["median_mbps"] / direct["median_mbps"], }, + "controller_proof": { + "network_profile": { + "one_way_delay_ms": int(delay_ms), + "loss_each_direction": loss, + "rate_each_direction": rate, + }, + "bbr": { + "artifact": bbr_path.name, + "acceleration": bbr["acceleration"], + "negotiated_tx_bytes_per_second": bbr.get( + "negotiated_tx_bytes_per_second", 0 + ), + "median_mbps": bbr["median_mbps"], + }, + "reno": { + "artifact": reno_path.name, + "acceleration": reno["acceleration"], + "negotiated_tx_bytes_per_second": reno.get( + "negotiated_tx_bytes_per_second", 0 + ), + "median_mbps": reno["median_mbps"], + }, + "bbr_to_reno_ratio": bbr_reno_ratio, + "relay_sender_download": { + "bbr_artifact": bbr_download_path.name, + "reno_artifact": reno_download_path.name, + "bbr_median_mbps": bbr_download["median_mbps"], + "reno_median_mbps": reno_download["median_mbps"], + "bbr_to_reno_ratio": bbr_reno_download_ratio, + "minimum_accepted_bbr_to_reno_ratio": float(minimum_bbr_reno_ratio), + }, + "minimum_accepted_bbr_to_reno_ratio": float(minimum_bbr_reno_ratio), + "brutal": { + "artifact": brutal_path.name, + "acceleration": brutal["acceleration"], + "negotiated_tx_bytes_per_second": brutal[ + "negotiated_tx_bytes_per_second" + ], + "median_mbps": brutal["median_mbps"], + "target_mbps": brutal_target_mbps, + "achieved_to_target_ratio": brutal_target_ratio, + "minimum_accepted_target_ratio": float(minimum_brutal_target_ratio), + "client_upload_mbps": int(client_upload_mbps), + "client_download_mbps": int(client_download_mbps), + "server_upload_cap_mbps": int(server_upload_mbps), + "server_download_cap_mbps": int(server_download_mbps), + }, + }, "acceleration_profile": { "description": "sequential short downloads over a loss-free high-RTT path; QUIC uses declared warmups", "payload_bytes": int(short_flow_bytes), diff --git a/third_party/hysteria-core/AUTOCAR_PATCHES.md b/third_party/hysteria-core/AUTOCAR_PATCHES.md new file mode 100644 index 0000000..0048c1a --- /dev/null +++ b/third_party/hysteria-core/AUTOCAR_PATCHES.md @@ -0,0 +1,51 @@ +# AutoCAR security hardening + +This directory is based on `github.com/apernet/hysteria/core/v2` v2.12.1 +(upstream commit `14e9fff1d972ab0187ac7fcf75b9514dc8664065`) and remains licensed under +the MIT license in `LICENSE.md`. + +AutoCAR keeps the fork intentionally small and auditable. Its server adds: + +- process-wide and per-source (IPv4 address or IPv6 `/64`) caps on accepted QUIC connections after Retry + address validation and before handshake state; +- process-wide and per-source caps on active TCP handlers, shared across + QUIC connections and held for the complete relay lifetime, plus a deadline + for reading the initial TCP request; +- a finite pre-authentication lifetime and a reduced incoming unidirectional + stream budget for unauthenticated HTTP/3 peers; +- process-wide and per-source UDP session admission, shared across QUIC + connections, before allocating defragmentation state; and +- fixed UDP fragment-count and reassembled-size bounds. + +Maximum-size UDP payloads use a framing-aware serialization buffer and every +serialization/fragmentation overflow returns an error instead of reporting a +successful silent drop. +The client package exports the 4,096-byte logical payload ceiling so frontends +can reject larger payloads without destroying a healthy UDP association. + +Authentication state and its identity are read under the same lock used by the +HTTP authentication handler, preventing dispatch or disconnect accounting from +observing a partially published authentication result. + +The QUIC listener forces Retry address validation and acquires the global +connection budget before allocating handshake state. Unauthenticated HTTP/3 +request headers are capped at 16 KiB before allocation. + +All per-source budgets use one IPv4 address or a masked IPv6 `/64`, preventing +interface-identifier rotation from bypassing the gates while documenting the +intentional NAT/prefix sharing tradeoff. + +These changes close resource-exhaustion paths that cannot be intercepted by +the public `server.Outbound` API because incomplete fragments never create an +outbound socket. The TCP admission ordering relies on the narrow local +`../quic-go` `StreamAdmission` hook documented in its `AUTOCAR_PATCHES.md`. +Changes should be rebased and re-audited whenever either pinned upstream +version changes. + +The upstream multi-gigabyte TCP and long lossy-UDP stress cases are opt-in via +`AUTOCAR_RUN_UPSTREAM_STRESS=1`; normal CI runs deterministic integration and +network-emulation suites instead of allowing those unbounded cases to consume +the job timeout. + +Integration tests generate an ephemeral ECDSA P-256 certificate in memory; +the fork does not carry the upstream repository's fixed test private key. diff --git a/third_party/hysteria-core/LICENSE.md b/third_party/hysteria-core/LICENSE.md new file mode 100644 index 0000000..208e8f2 --- /dev/null +++ b/third_party/hysteria-core/LICENSE.md @@ -0,0 +1,7 @@ +Copyright 2023 Toby + +Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the “Software”), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED “AS IS”, WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. diff --git a/third_party/hysteria-core/client/.mockery.yaml b/third_party/hysteria-core/client/.mockery.yaml new file mode 100644 index 0000000..299e6f9 --- /dev/null +++ b/third_party/hysteria-core/client/.mockery.yaml @@ -0,0 +1,9 @@ +with-expecter: true +inpackage: true +dir: . +packages: + github.com/apernet/hysteria/core/v2/client: + interfaces: + udpIO: + config: + mockname: mockUDPIO diff --git a/third_party/hysteria-core/client/client.go b/third_party/hysteria-core/client/client.go new file mode 100644 index 0000000..fe46ee4 --- /dev/null +++ b/third_party/hysteria-core/client/client.go @@ -0,0 +1,383 @@ +package client + +import ( + "context" + "crypto/tls" + "errors" + "net" + "net/http" + "net/url" + "sync" + "time" + + coreErrs "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/congestion" + "github.com/apernet/hysteria/core/v2/internal/protocol" + "github.com/apernet/hysteria/core/v2/internal/utils" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" +) + +const ( + closeErrCodeOK = 0x100 // HTTP3 ErrCodeNoError + closeErrCodeProtocolError = 0x101 // HTTP3 ErrCodeGeneralProtocolError + + // MaxUDPSize is the largest logical UDP payload carried by one Hysteria + // message. Larger messages must be rejected before serialization. + MaxUDPSize = protocol.MaxUDPSize +) + +type Client interface { + TCP(addr string) (net.Conn, error) + UDP() (HyUDPConn, error) + Close() error +} + +type HyUDPConn interface { + Receive() ([]byte, string, error) + Send([]byte, string) error + Close() error +} + +type HandshakeInfo struct { + UDPEnabled bool + Tx uint64 // 0 if using BBR + ServerAddr net.Addr + ECHAccepted bool +} + +func NewClient(config *Config) (Client, *HandshakeInfo, error) { + if err := config.verifyAndFill(); err != nil { + return nil, nil, err + } + c := &clientImpl{ + config: config, + } + info, err := c.connect() + if err != nil { + return nil, nil, err + } + return c, info, nil +} + +type clientImpl struct { + config *Config + + pktConn net.PacketConn + tr *quic.Transport + conn *quic.Conn + + udpSM *udpSessionManager +} + +func (c *clientImpl) connect() (*HandshakeInfo, error) { + pktConn, err := c.config.ConnFactory.New(c.config.ServerAddr) + if err != nil { + return nil, err + } + // Convert config to TLS config & QUIC config + tlsConfig := &tls.Config{ + ServerName: c.config.TLSConfig.ServerName, + InsecureSkipVerify: c.config.TLSConfig.InsecureSkipVerify, + VerifyPeerCertificate: c.config.TLSConfig.VerifyPeerCertificate, + RootCAs: c.config.TLSConfig.RootCAs, + GetClientCertificate: c.config.TLSConfig.GetClientCertificate, + EncryptedClientHelloConfigList: c.config.TLSConfig.ECHConfigList, + } + quicConfig := &quic.Config{ + InitialStreamReceiveWindow: c.config.QUICConfig.InitialStreamReceiveWindow, + MaxStreamReceiveWindow: c.config.QUICConfig.MaxStreamReceiveWindow, + InitialConnectionReceiveWindow: c.config.QUICConfig.InitialConnectionReceiveWindow, + MaxConnectionReceiveWindow: c.config.QUICConfig.MaxConnectionReceiveWindow, + MaxIdleTimeout: c.config.QUICConfig.MaxIdleTimeout, + KeepAlivePeriod: c.config.QUICConfig.KeepAlivePeriod, + DisablePathMTUDiscovery: c.config.QUICConfig.DisablePathMTUDiscovery, + EnableDatagrams: true, + MaxDatagramFrameSize: protocol.MaxDatagramFrameSize, + OmitMaxDatagramFrameSize: true, + DisablePathManager: true, + ChromeParrot: !c.config.QUICConfig.DisableChromeParrot, + } + tr := &quic.Transport{Conn: pktConn, DisableGSO: c.config.QUICConfig.DisableGSO} + if !c.config.QUICConfig.DisableChromeParrot { + // Chrome uses a zero-length source connection ID. This has to be set on the + // Transport, since it fixes the length at which incoming packets' connection + // IDs are parsed; leaving it default yields 4-byte IDs, visible on the wire. + tr.ConnectionIDGenerator = quic.ZeroLengthConnectionIDGenerator{} + } + // Prepare RoundTripper + var conn *quic.Conn + rt := &http3.Transport{ + TLSClientConfig: tlsConfig, + QUICConfig: quicConfig, + Dial: func(ctx context.Context, _ string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + qc, err := tr.DialEarly(ctx, c.config.ServerAddr, tlsCfg, cfg) + if err != nil { + return nil, err + } + conn = qc + return qc, nil + }, + } + // Send auth HTTP request + req := &http.Request{ + Method: http.MethodPost, + URL: &url.URL{ + Scheme: "https", + Host: protocol.URLHost, + Path: protocol.URLPath, + }, + Header: make(http.Header), + } + protocol.AuthRequestToHeader(req.Header, protocol.AuthRequest{ + Auth: c.config.Auth, + Rx: c.config.BandwidthConfig.MaxRx, + }) + resp, err := rt.RoundTrip(req) + if err != nil { + if conn != nil { + _ = conn.CloseWithError(closeErrCodeProtocolError, "") + } + _ = tr.Close() + _ = pktConn.Close() + return nil, coreErrs.ConnectError{Err: err} + } + if resp.StatusCode != protocol.StatusAuthOK { + _ = conn.CloseWithError(closeErrCodeProtocolError, "") + _ = tr.Close() + _ = pktConn.Close() + return nil, coreErrs.AuthError{StatusCode: resp.StatusCode} + } + // Auth OK + authResp := protocol.AuthResponseFromHeader(resp.Header) + var actualTx uint64 + if authResp.RxAuto { + // Server asks client to use bandwidth detection, + // ignore local bandwidth config and use the configured congestion controller. + congestion.UseConfigured(conn, c.config.CongestionConfig.Type, c.config.CongestionConfig.BBRProfile) + } else { + // actualTx = min(serverRx, clientTx) + actualTx = authResp.Rx + if actualTx == 0 || actualTx > c.config.BandwidthConfig.MaxTx { + // Server doesn't have a limit, or our clientTx is smaller than serverRx + actualTx = c.config.BandwidthConfig.MaxTx + } + if actualTx > 0 { + congestion.UseBrutal(conn, actualTx, c.config.BandwidthConfig.DisableLossCompensation) + } else { + // We don't know our own bandwidth either, use the configured congestion controller. + congestion.UseConfigured(conn, c.config.CongestionConfig.Type, c.config.CongestionConfig.BBRProfile) + } + } + _ = resp.Body.Close() + + c.pktConn = pktConn + c.tr = tr + c.conn = conn + if authResp.UDPEnabled { + c.udpSM = newUDPSessionManager(&udpIOImpl{Conn: conn}) + } + return &HandshakeInfo{ + UDPEnabled: authResp.UDPEnabled, + Tx: actualTx, + ServerAddr: c.config.ServerAddr, + ECHAccepted: conn.ConnectionState().TLS.ECHAccepted, + }, nil +} + +// openStream wraps the stream with QStream, which handles Close() properly +func (c *clientImpl) openStream() (*utils.QStream, error) { + stream, err := c.conn.OpenStream() + if err != nil { + return nil, err + } + return &utils.QStream{Stream: stream}, nil +} + +func (c *clientImpl) TCP(addr string) (net.Conn, error) { + stream, err := c.openStream() + if err != nil { + return nil, wrapIfConnectionClosed(err) + } + // Send request + err = protocol.WriteTCPRequest(stream, addr) + if err != nil { + _ = stream.Close() + return nil, wrapIfConnectionClosed(err) + } + if c.config.FastOpen { + // Don't wait for the response when fast open is enabled. + // Return the connection immediately, defer the response handling + // to the first Read() call. + return &tcpConn{ + Orig: stream, + PseudoLocalAddr: c.conn.LocalAddr(), + PseudoRemoteAddr: c.conn.RemoteAddr(), + }, nil + } + // Read response + ok, msg, err := protocol.ReadTCPResponse(stream) + if err != nil { + _ = stream.Close() + return nil, wrapIfConnectionClosed(err) + } + if !ok { + _ = stream.Close() + return nil, coreErrs.DialError{Message: msg} + } + return &tcpConn{ + Orig: stream, + PseudoLocalAddr: c.conn.LocalAddr(), + PseudoRemoteAddr: c.conn.RemoteAddr(), + established: true, + }, nil +} + +func (c *clientImpl) UDP() (HyUDPConn, error) { + if c.udpSM == nil { + return nil, coreErrs.DialError{Message: "UDP not enabled"} + } + return c.udpSM.NewUDP() +} + +func (c *clientImpl) Close() error { + _ = c.conn.CloseWithError(closeErrCodeOK, "") + _ = c.tr.Close() + _ = c.pktConn.Close() + return nil +} + +var nonPermanentErrors = []error{ + quic.StreamLimitReachedError{}, +} + +// wrapIfConnectionClosed checks if the error returned by quic-go +// is recoverable (listed in nonPermanentErrors) or permanent. +// Recoverable errors are returned as-is, +// permanent ones are wrapped as ClosedError. +func wrapIfConnectionClosed(err error) error { + for _, e := range nonPermanentErrors { + if errors.Is(err, e) { + return err + } + } + return coreErrs.ClosedError{Err: err} +} + +type tcpStream interface { + Read([]byte) (int, error) + Write([]byte) (int, error) + Close() error + SetDeadline(time.Time) error + SetReadDeadline(time.Time) error + SetWriteDeadline(time.Time) error +} + +type tcpConn struct { + Orig tcpStream + PseudoLocalAddr net.Addr + PseudoRemoteAddr net.Addr + + establishMu sync.Mutex + established bool + establishErr error + closeMu sync.Mutex + closed bool + closeErr error +} + +func (c *tcpConn) Read(b []byte) (n int, err error) { + if err := c.ensureEstablished(); err != nil { + return 0, err + } + return c.Orig.Read(b) +} + +func (c *tcpConn) ensureEstablished() error { + c.establishMu.Lock() + defer c.establishMu.Unlock() + if c.established { + return nil + } + if c.establishErr != nil { + return c.establishErr + } + ok, msg, err := protocol.ReadTCPResponse(c.Orig) + if err != nil { + c.establishErr = err + } else if !ok { + c.establishErr = coreErrs.DialError{Message: msg} + } else { + c.established = true + return nil + } + _ = c.closeOrig() + return c.establishErr +} + +func (c *tcpConn) Write(b []byte) (n int, err error) { + return c.Orig.Write(b) +} + +func (c *tcpConn) Close() error { + return c.closeOrig() +} + +func (c *tcpConn) closeOrig() error { + c.closeMu.Lock() + defer c.closeMu.Unlock() + if !c.closed { + c.closeErr = c.Orig.Close() + c.closed = true + } + return c.closeErr +} + +func (c *tcpConn) LocalAddr() net.Addr { + return c.PseudoLocalAddr +} + +func (c *tcpConn) RemoteAddr() net.Addr { + return c.PseudoRemoteAddr +} + +func (c *tcpConn) SetDeadline(t time.Time) error { + return c.Orig.SetDeadline(t) +} + +func (c *tcpConn) SetReadDeadline(t time.Time) error { + return c.Orig.SetReadDeadline(t) +} + +func (c *tcpConn) SetWriteDeadline(t time.Time) error { + return c.Orig.SetWriteDeadline(t) +} + +type udpIOImpl struct { + Conn *quic.Conn +} + +func (io *udpIOImpl) ReceiveMessage() (*protocol.UDPMessage, error) { + for { + msg, err := io.Conn.ReceiveDatagram(context.Background()) + if err != nil { + // Connection error, this will stop the session manager + return nil, err + } + udpMsg, err := protocol.ParseUDPMessage(msg) + if err != nil { + // Invalid message, this is fine - just wait for the next + continue + } + return udpMsg, nil + } +} + +func (io *udpIOImpl) SendMessage(buf []byte, msg *protocol.UDPMessage) error { + msgN := msg.Serialize(buf) + if msgN < 0 { + return coreErrs.ProtocolError{Message: "UDP message exceeds serialization limit"} + } + return io.Conn.SendDatagram(buf[:msgN]) +} diff --git a/third_party/hysteria-core/client/config.go b/third_party/hysteria-core/client/config.go new file mode 100644 index 0000000..131a213 --- /dev/null +++ b/third_party/hysteria-core/client/config.go @@ -0,0 +1,136 @@ +package client + +import ( + "crypto/tls" + "crypto/x509" + "net" + "time" + + "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/congestion" + "github.com/apernet/hysteria/core/v2/internal/pmtud" +) + +const ( + defaultStreamReceiveWindow = 8388608 // 8MB + defaultConnReceiveWindow = defaultStreamReceiveWindow * 5 / 2 // 20MB + defaultMaxIdleTimeout = 30 * time.Second + defaultKeepAlivePeriod = 10 * time.Second +) + +type Config struct { + ConnFactory ConnFactory + ServerAddr net.Addr + Auth string + TLSConfig TLSConfig + QUICConfig QUICConfig + CongestionConfig CongestionConfig + BandwidthConfig BandwidthConfig + FastOpen bool + + filled bool // whether the fields have been verified and filled +} + +// verifyAndFill fills the fields that are not set by the user with default values when possible, +// and returns an error if the user has not set a required field or has set an invalid value. +func (c *Config) verifyAndFill() error { + if c.filled { + return nil + } + if c.ConnFactory == nil { + c.ConnFactory = &udpConnFactory{} + } + if c.ServerAddr == nil { + return errors.ConfigError{Field: "ServerAddr", Reason: "must be set"} + } + if c.QUICConfig.InitialStreamReceiveWindow == 0 { + c.QUICConfig.InitialStreamReceiveWindow = defaultStreamReceiveWindow + } else if c.QUICConfig.InitialStreamReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.InitialStreamReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.MaxStreamReceiveWindow == 0 { + c.QUICConfig.MaxStreamReceiveWindow = defaultStreamReceiveWindow + } else if c.QUICConfig.MaxStreamReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.MaxStreamReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.InitialConnectionReceiveWindow == 0 { + c.QUICConfig.InitialConnectionReceiveWindow = defaultConnReceiveWindow + } else if c.QUICConfig.InitialConnectionReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.InitialConnectionReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.MaxConnectionReceiveWindow == 0 { + c.QUICConfig.MaxConnectionReceiveWindow = defaultConnReceiveWindow + } else if c.QUICConfig.MaxConnectionReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.MaxConnectionReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.MaxIdleTimeout == 0 { + c.QUICConfig.MaxIdleTimeout = defaultMaxIdleTimeout + } else if c.QUICConfig.MaxIdleTimeout < 4*time.Second || c.QUICConfig.MaxIdleTimeout > 120*time.Second { + return errors.ConfigError{Field: "QUICConfig.MaxIdleTimeout", Reason: "must be between 4s and 120s"} + } + if c.QUICConfig.KeepAlivePeriod == 0 { + c.QUICConfig.KeepAlivePeriod = defaultKeepAlivePeriod + } else if c.QUICConfig.KeepAlivePeriod < 2*time.Second || c.QUICConfig.KeepAlivePeriod > 60*time.Second { + return errors.ConfigError{Field: "QUICConfig.KeepAlivePeriod", Reason: "must be between 2s and 60s"} + } + c.QUICConfig.DisablePathMTUDiscovery = c.QUICConfig.DisablePathMTUDiscovery || pmtud.DisablePathMTUDiscovery + var err error + c.CongestionConfig.Type, err = congestion.NormalizeType(c.CongestionConfig.Type) + if err != nil { + return errors.ConfigError{Field: "CongestionConfig.Type", Reason: err.Error()} + } + if c.CongestionConfig.Type == congestion.TypeBBR { + c.CongestionConfig.BBRProfile, err = congestion.NormalizeBBRProfile(c.CongestionConfig.BBRProfile) + if err != nil { + return errors.ConfigError{Field: "CongestionConfig.BBRProfile", Reason: err.Error()} + } + } + + c.filled = true + return nil +} + +type ConnFactory interface { + New(net.Addr) (net.PacketConn, error) +} + +type udpConnFactory struct{} + +func (f *udpConnFactory) New(addr net.Addr) (net.PacketConn, error) { + return net.ListenUDP("udp", nil) +} + +// TLSConfig contains the TLS configuration fields that we want to expose to the user. +type TLSConfig struct { + ServerName string + InsecureSkipVerify bool + VerifyPeerCertificate func(rawCerts [][]byte, verifiedChains [][]*x509.Certificate) error + RootCAs *x509.CertPool + GetClientCertificate func(*tls.CertificateRequestInfo) (*tls.Certificate, error) + ECHConfigList []byte +} + +// QUICConfig contains the QUIC configuration fields that we want to expose to the user. +type QUICConfig struct { + InitialStreamReceiveWindow uint64 + MaxStreamReceiveWindow uint64 + InitialConnectionReceiveWindow uint64 + MaxConnectionReceiveWindow uint64 + MaxIdleTimeout time.Duration + KeepAlivePeriod time.Duration + DisablePathMTUDiscovery bool // The server may still override this to true on unsupported platforms. + DisableGSO bool + DisableChromeParrot bool // Chrome QUIC fingerprint parroting is on by default. +} + +type CongestionConfig struct { + Type string + BBRProfile string +} + +// BandwidthConfig describes the maximum bandwidth that the server can use, in bytes per second. +type BandwidthConfig struct { + MaxTx uint64 + MaxRx uint64 + DisableLossCompensation bool +} diff --git a/third_party/hysteria-core/client/fast_open_test.go b/third_party/hysteria-core/client/fast_open_test.go new file mode 100644 index 0000000..4f56dd3 --- /dev/null +++ b/third_party/hysteria-core/client/fast_open_test.go @@ -0,0 +1,87 @@ +package client + +import ( + "errors" + "net" + "sort" + "sync" + "testing" + "time" + + coreErrs "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +func TestFastOpenConcurrentReadsConsumeResponseOnce(t *testing.T) { + clientSide, serverSide := net.Pipe() + defer serverSide.Close() + conn := &tcpConn{Orig: clientSide} + go func() { + _ = protocol.WriteTCPResponse(serverSide, true, "connected") + _, _ = serverSide.Write([]byte("xy")) + }() + + start := make(chan struct{}) + results := make(chan byte, 2) + errorsSeen := make(chan error, 2) + var wait sync.WaitGroup + for range 2 { + wait.Add(1) + go func() { + defer wait.Done() + <-start + buffer := make([]byte, 1) + _, err := conn.Read(buffer) + if err != nil { + errorsSeen <- err + return + } + results <- buffer[0] + }() + } + close(start) + wait.Wait() + close(errorsSeen) + for err := range errorsSeen { + t.Fatalf("concurrent Read: %v", err) + } + close(results) + var got []byte + for result := range results { + got = append(got, result) + } + sort.Slice(got, func(i, j int) bool { return got[i] < got[j] }) + if string(got) != "xy" { + t.Fatalf("concurrent payload = %q, want xy", got) + } +} + +func TestFastOpenFailureIsCachedAndClosesStream(t *testing.T) { + clientSide, serverSide := net.Pipe() + conn := &tcpConn{Orig: clientSide} + go func() { + _ = protocol.WriteTCPResponse(serverSide, false, "destination denied") + _ = serverSide.Close() + }() + + buffer := make([]byte, 1) + _, firstErr := conn.Read(buffer) + var firstDialError coreErrs.DialError + if !errors.As(firstErr, &firstDialError) { + t.Fatalf("first Read error = %v, want DialError", firstErr) + } + done := make(chan error, 1) + go func() { + _, err := conn.Read(buffer) + done <- err + }() + select { + case secondErr := <-done: + var secondDialError coreErrs.DialError + if !errors.As(secondErr, &secondDialError) || secondErr.Error() != firstErr.Error() { + t.Fatalf("cached Read error = %v, want %v", secondErr, firstErr) + } + case <-time.After(time.Second): + t.Fatal("second Read retried the consumed Fast Open response") + } +} diff --git a/third_party/hysteria-core/client/mock_udpIO.go b/third_party/hysteria-core/client/mock_udpIO.go new file mode 100644 index 0000000..dbff53c --- /dev/null +++ b/third_party/hysteria-core/client/mock_udpIO.go @@ -0,0 +1,139 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package client + +import ( + protocol "github.com/apernet/hysteria/core/v2/internal/protocol" + mock "github.com/stretchr/testify/mock" +) + +// mockUDPIO is an autogenerated mock type for the udpIO type +type mockUDPIO struct { + mock.Mock +} + +type mockUDPIO_Expecter struct { + mock *mock.Mock +} + +func (_m *mockUDPIO) EXPECT() *mockUDPIO_Expecter { + return &mockUDPIO_Expecter{mock: &_m.Mock} +} + +// ReceiveMessage provides a mock function with no fields +func (_m *mockUDPIO) ReceiveMessage() (*protocol.UDPMessage, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for ReceiveMessage") + } + + var r0 *protocol.UDPMessage + var r1 error + if rf, ok := ret.Get(0).(func() (*protocol.UDPMessage, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() *protocol.UDPMessage); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*protocol.UDPMessage) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// mockUDPIO_ReceiveMessage_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ReceiveMessage' +type mockUDPIO_ReceiveMessage_Call struct { + *mock.Call +} + +// ReceiveMessage is a helper method to define mock.On call +func (_e *mockUDPIO_Expecter) ReceiveMessage() *mockUDPIO_ReceiveMessage_Call { + return &mockUDPIO_ReceiveMessage_Call{Call: _e.mock.On("ReceiveMessage")} +} + +func (_c *mockUDPIO_ReceiveMessage_Call) Run(run func()) *mockUDPIO_ReceiveMessage_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *mockUDPIO_ReceiveMessage_Call) Return(_a0 *protocol.UDPMessage, _a1 error) *mockUDPIO_ReceiveMessage_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *mockUDPIO_ReceiveMessage_Call) RunAndReturn(run func() (*protocol.UDPMessage, error)) *mockUDPIO_ReceiveMessage_Call { + _c.Call.Return(run) + return _c +} + +// SendMessage provides a mock function with given fields: _a0, _a1 +func (_m *mockUDPIO) SendMessage(_a0 []byte, _a1 *protocol.UDPMessage) error { + ret := _m.Called(_a0, _a1) + + if len(ret) == 0 { + panic("no return value specified for SendMessage") + } + + var r0 error + if rf, ok := ret.Get(0).(func([]byte, *protocol.UDPMessage) error); ok { + r0 = rf(_a0, _a1) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockUDPIO_SendMessage_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendMessage' +type mockUDPIO_SendMessage_Call struct { + *mock.Call +} + +// SendMessage is a helper method to define mock.On call +// - _a0 []byte +// - _a1 *protocol.UDPMessage +func (_e *mockUDPIO_Expecter) SendMessage(_a0 interface{}, _a1 interface{}) *mockUDPIO_SendMessage_Call { + return &mockUDPIO_SendMessage_Call{Call: _e.mock.On("SendMessage", _a0, _a1)} +} + +func (_c *mockUDPIO_SendMessage_Call) Run(run func(_a0 []byte, _a1 *protocol.UDPMessage)) *mockUDPIO_SendMessage_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte), args[1].(*protocol.UDPMessage)) + }) + return _c +} + +func (_c *mockUDPIO_SendMessage_Call) Return(_a0 error) *mockUDPIO_SendMessage_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockUDPIO_SendMessage_Call) RunAndReturn(run func([]byte, *protocol.UDPMessage) error) *mockUDPIO_SendMessage_Call { + _c.Call.Return(run) + return _c +} + +// newMockUDPIO creates a new instance of mockUDPIO. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func newMockUDPIO(t interface { + mock.TestingT + Cleanup(func()) +}) *mockUDPIO { + mock := &mockUDPIO{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/client/reconnect.go b/third_party/hysteria-core/client/reconnect.go new file mode 100644 index 0000000..91c7eb4 --- /dev/null +++ b/third_party/hysteria-core/client/reconnect.go @@ -0,0 +1,120 @@ +package client + +import ( + "net" + "sync" + + coreErrs "github.com/apernet/hysteria/core/v2/errors" +) + +// reconnectableClientImpl is a wrapper of Client, which can reconnect when the connection is closed, +// except when the caller explicitly calls Close() to permanently close this client. +type reconnectableClientImpl struct { + configFunc func() (*Config, error) // called before connecting + connectedFunc func(Client, *HandshakeInfo, int) // called when successfully connected + client Client + count int + m sync.Mutex + closed bool // permanent close +} + +// NewReconnectableClient creates a reconnectable client. +// If lazy is true, the client will not connect until the first call to TCP() or UDP(). +// We use a function for config mainly to delay config evaluation +// (which involves DNS resolution) until the actual connection attempt. +func NewReconnectableClient(configFunc func() (*Config, error), connectedFunc func(Client, *HandshakeInfo, int), lazy bool) (Client, error) { + rc := &reconnectableClientImpl{ + configFunc: configFunc, + connectedFunc: connectedFunc, + } + if !lazy { + if err := rc.reconnect(); err != nil { + return nil, err + } + } + return rc, nil +} + +func (rc *reconnectableClientImpl) reconnect() error { + if rc.client != nil { + _ = rc.client.Close() + } + var info *HandshakeInfo + config, err := rc.configFunc() + if err != nil { + return err + } + rc.client, info, err = NewClient(config) + if err != nil { + return err + } else { + rc.count++ + if rc.connectedFunc != nil { + rc.connectedFunc(rc, info, rc.count) + } + return nil + } +} + +// clientDo calls f with the current client. +// If the client is nil, it will first reconnect. +// It will also detect if the client is closed, and if so, +// set it to nil for reconnect next time. +func (rc *reconnectableClientImpl) clientDo(f func(Client) (interface{}, error)) (interface{}, error) { + rc.m.Lock() + if rc.closed { + rc.m.Unlock() + return nil, coreErrs.ClosedError{} + } + if rc.client == nil { + // No active connection, connect first + if err := rc.reconnect(); err != nil { + rc.m.Unlock() + return nil, err + } + } + client := rc.client + rc.m.Unlock() + + ret, err := f(client) + if _, ok := err.(coreErrs.ClosedError); ok { + // Connection closed, set client to nil for reconnect next time + rc.m.Lock() + if rc.client == client { + // This check is in case the client is already changed by another goroutine + rc.client = nil + } + rc.m.Unlock() + } + return ret, err +} + +func (rc *reconnectableClientImpl) TCP(addr string) (net.Conn, error) { + if c, err := rc.clientDo(func(client Client) (interface{}, error) { + return client.TCP(addr) + }); err != nil { + return nil, err + } else { + return c.(net.Conn), nil + } +} + +func (rc *reconnectableClientImpl) UDP() (HyUDPConn, error) { + if c, err := rc.clientDo(func(client Client) (interface{}, error) { + return client.UDP() + }); err != nil { + return nil, err + } else { + return c.(HyUDPConn), nil + } +} + +func (rc *reconnectableClientImpl) Close() error { + rc.m.Lock() + defer rc.m.Unlock() + rc.closed = true + if rc.client != nil { + return rc.client.Close() + } + return nil +} diff --git a/third_party/hysteria-core/client/udp.go b/third_party/hysteria-core/client/udp.go new file mode 100644 index 0000000..c2e7a30 --- /dev/null +++ b/third_party/hysteria-core/client/udp.go @@ -0,0 +1,188 @@ +package client + +import ( + "errors" + "io" + "math/rand" + "sync" + + "github.com/apernet/quic-go" + + coreErrs "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/frag" + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +const ( + udpMessageChanSize = 1024 +) + +type udpIO interface { + ReceiveMessage() (*protocol.UDPMessage, error) + SendMessage([]byte, *protocol.UDPMessage) error +} + +type udpConn struct { + ID uint32 + D *frag.Defragger + ReceiveCh chan *protocol.UDPMessage + SendBuf []byte + SendFunc func([]byte, *protocol.UDPMessage) error + CloseFunc func() + Closed bool +} + +func (u *udpConn) Receive() ([]byte, string, error) { + for { + msg := <-u.ReceiveCh + if msg == nil { + // Closed + return nil, "", io.EOF + } + dfMsg := u.D.Feed(msg) + if dfMsg == nil { + // Incomplete message, wait for more + continue + } + return dfMsg.Data, dfMsg.Addr, nil + } +} + +// Send is not thread-safe, as it uses a shared SendBuf. +func (u *udpConn) Send(data []byte, addr string) error { + // Try no frag first + msg := &protocol.UDPMessage{ + SessionID: u.ID, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: addr, + Data: data, + } + err := u.SendFunc(u.SendBuf, msg) + var errTooLarge *quic.DatagramTooLargeError + if errors.As(err, &errTooLarge) { + // Message too large, try fragmentation + msg.PacketID = uint16(rand.Intn(0xFFFF)) + 1 + fMsgs := frag.FragUDPMessage(msg, int(errTooLarge.MaxDatagramPayloadSize)) + if len(fMsgs) == 0 { + return frag.ErrFragmentationLimit + } + for _, fMsg := range fMsgs { + err := u.SendFunc(u.SendBuf, &fMsg) + if err != nil { + return err + } + } + return nil + } else { + return err + } +} + +func (u *udpConn) Close() error { + u.CloseFunc() + return nil +} + +type udpSessionManager struct { + io udpIO + + mutex sync.RWMutex + m map[uint32]*udpConn + nextID uint32 + + closed bool +} + +func newUDPSessionManager(io udpIO) *udpSessionManager { + m := &udpSessionManager{ + io: io, + m: make(map[uint32]*udpConn), + nextID: 1, + } + go m.run() + return m +} + +func (m *udpSessionManager) run() error { + defer m.closeCleanup() + for { + msg, err := m.io.ReceiveMessage() + if err != nil { + return err + } + m.feed(msg) + } +} + +func (m *udpSessionManager) closeCleanup() { + m.mutex.Lock() + defer m.mutex.Unlock() + + for _, conn := range m.m { + m.close(conn) + } + m.closed = true +} + +func (m *udpSessionManager) feed(msg *protocol.UDPMessage) { + m.mutex.RLock() + defer m.mutex.RUnlock() + + conn, ok := m.m[msg.SessionID] + if !ok { + // Ignore message from unknown session + return + } + + select { + case conn.ReceiveCh <- msg: + // OK + default: + // Channel full, drop the message + } +} + +// NewUDP creates a new UDP session. +func (m *udpSessionManager) NewUDP() (HyUDPConn, error) { + m.mutex.Lock() + defer m.mutex.Unlock() + + if m.closed { + return nil, coreErrs.ClosedError{} + } + + id := m.nextID + m.nextID++ + + conn := &udpConn{ + ID: id, + D: &frag.Defragger{}, + ReceiveCh: make(chan *protocol.UDPMessage, udpMessageChanSize), + SendBuf: make([]byte, protocol.MaxUDPMessageSize), + SendFunc: m.io.SendMessage, + } + conn.CloseFunc = func() { + m.mutex.Lock() + defer m.mutex.Unlock() + m.close(conn) + } + m.m[id] = conn + + return conn, nil +} + +func (m *udpSessionManager) close(conn *udpConn) { + if !conn.Closed { + conn.Closed = true + close(conn.ReceiveCh) + delete(m.m, conn.ID) + } +} + +func (m *udpSessionManager) Count() int { + m.mutex.RLock() + defer m.mutex.RUnlock() + return len(m.m) +} diff --git a/third_party/hysteria-core/client/udp_test.go b/third_party/hysteria-core/client/udp_test.go new file mode 100644 index 0000000..d039b5b --- /dev/null +++ b/third_party/hysteria-core/client/udp_test.go @@ -0,0 +1,146 @@ +package client + +import ( + "errors" + io2 "io" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.uber.org/goleak" + + coreErrs "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/frag" + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +func TestUDPSendReportsFragmentationLimit(t *testing.T) { + conn := &udpConn{ + ID: 1, + SendBuf: make([]byte, protocol.MaxUDPSize), + SendFunc: func([]byte, *protocol.UDPMessage) error { + return &quic.DatagramTooLargeError{MaxDatagramPayloadSize: 1} + }, + } + err := conn.Send(make([]byte, protocol.MaxUDPSize), "example.test:443") + if !errors.Is(err, frag.ErrFragmentationLimit) { + t.Fatalf("Send error = %v, want fragmentation limit", err) + } +} + +func TestUDPSerializationOverflowIsReported(t *testing.T) { + io := &udpIOImpl{} + err := io.SendMessage(make([]byte, 1), &protocol.UDPMessage{Addr: "example.test:443", Data: []byte("payload")}) + if err == nil { + t.Fatal("serialization overflow was silently dropped") + } +} + +func TestUDPSessionManager(t *testing.T) { + io := newMockUDPIO(t) + receiveCh := make(chan *protocol.UDPMessage, 4) + io.EXPECT().ReceiveMessage().RunAndReturn(func() (*protocol.UDPMessage, error) { + m := <-receiveCh + if m == nil { + return nil, errors.New("closed") + } + return m, nil + }) + sm := newUDPSessionManager(io) + + // Test UDP session IO + udpConn1, err := sm.NewUDP() + assert.NoError(t, err) + udpConn2, err := sm.NewUDP() + assert.NoError(t, err) + + msg1 := &protocol.UDPMessage{ + SessionID: 1, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "random.site.com:9000", + Data: []byte("hello friend"), + } + io.EXPECT().SendMessage(mock.Anything, msg1).Return(nil).Once() + err = udpConn1.Send(msg1.Data, msg1.Addr) + assert.NoError(t, err) + + msg2 := &protocol.UDPMessage{ + SessionID: 2, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "another.site.org:8000", + Data: []byte("mr robot"), + } + io.EXPECT().SendMessage(mock.Anything, msg2).Return(nil).Once() + err = udpConn2.Send(msg2.Data, msg2.Addr) + assert.NoError(t, err) + + respMsg1 := &protocol.UDPMessage{ + SessionID: 1, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: msg1.Addr, + Data: []byte("goodbye captain price"), + } + receiveCh <- respMsg1 + data, addr, err := udpConn1.Receive() + assert.NoError(t, err) + assert.Equal(t, data, respMsg1.Data) + assert.Equal(t, addr, respMsg1.Addr) + + respMsg2 := &protocol.UDPMessage{ + SessionID: 2, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: msg2.Addr, + Data: []byte("white rose"), + } + receiveCh <- respMsg2 + data, addr, err = udpConn2.Receive() + assert.NoError(t, err) + assert.Equal(t, data, respMsg2.Data) + assert.Equal(t, addr, respMsg2.Addr) + + respMsg3 := &protocol.UDPMessage{ + SessionID: 55, // Bogus session ID that doesn't exist + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "burgerking.com:27017", + Data: []byte("impossible whopper"), + } + receiveCh <- respMsg3 + // No test for this, just make sure it doesn't panic + + // Test close UDP connection unblocks Receive() + errChan := make(chan error, 1) + go func() { + _, _, err := udpConn1.Receive() + errChan <- err + }() + assert.NoError(t, udpConn1.Close()) + assert.Equal(t, <-errChan, io2.EOF) + + // Test close IO unblocks Receive() and blocks new UDP creation + errChan = make(chan error, 1) + go func() { + _, _, err := udpConn2.Receive() + errChan <- err + }() + close(receiveCh) + assert.Equal(t, <-errChan, io2.EOF) + _, err = sm.NewUDP() + assert.Equal(t, err, coreErrs.ClosedError{}) + + // Leak checks + time.Sleep(1 * time.Second) + assert.Zero(t, sm.Count(), "session count should be 0") + goleak.VerifyNone(t) +} diff --git a/third_party/hysteria-core/errors/errors.go b/third_party/hysteria-core/errors/errors.go new file mode 100644 index 0000000..cb69118 --- /dev/null +++ b/third_party/hysteria-core/errors/errors.go @@ -0,0 +1,75 @@ +package errors + +import ( + "fmt" + "strconv" +) + +// ConfigError is returned when a configuration field is invalid. +type ConfigError struct { + Field string + Reason string +} + +func (c ConfigError) Error() string { + return fmt.Sprintf("invalid config: %s: %s", c.Field, c.Reason) +} + +// ConnectError is returned when the client fails to connect to the server. +type ConnectError struct { + Err error +} + +func (c ConnectError) Error() string { + return "connect error: " + c.Err.Error() +} + +func (c ConnectError) Unwrap() error { + return c.Err +} + +// AuthError is returned when the client fails to authenticate with the server. +type AuthError struct { + StatusCode int +} + +func (a AuthError) Error() string { + return "authentication error, HTTP status code: " + strconv.Itoa(a.StatusCode) +} + +// DialError is returned when the server rejects the client's dial request. +// This applies to both TCP and UDP. +type DialError struct { + Message string +} + +func (c DialError) Error() string { + return "dial error: " + c.Message +} + +// ClosedError is returned when the client attempts to use a closed connection. +type ClosedError struct { + Err error // Can be nil +} + +func (c ClosedError) Error() string { + if c.Err == nil { + return "connection closed" + } else { + return "connection closed: " + c.Err.Error() + } +} + +func (c ClosedError) Unwrap() error { + return c.Err +} + +// ProtocolError is returned when the server/client runs into an unexpected +// or malformed request/response/message. +type ProtocolError struct { + Message string +} + +func (p ProtocolError) Error() string { + return "protocol error: " + p.Message +} diff --git a/third_party/hysteria-core/go.mod b/third_party/hysteria-core/go.mod new file mode 100644 index 0000000..98fbf11 --- /dev/null +++ b/third_party/hysteria-core/go.mod @@ -0,0 +1,32 @@ +module github.com/apernet/hysteria/core/v2 + +go 1.25.0 + +toolchain go1.25.1 + +replace github.com/apernet/quic-go => ../quic-go + +require ( + github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e + github.com/stretchr/testify v1.11.1 + go.uber.org/goleak v1.3.0 + golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 + golang.org/x/time v0.15.0 +) + +require ( + github.com/andybalholm/brotli v1.1.0 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/klauspost/compress v1.18.7 // indirect + github.com/kr/text v0.2.0 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/quic-go/qpack v0.6.0 // indirect + github.com/refraction-networking/utls v1.8.2 // indirect + github.com/rogpeppe/go-internal v1.12.0 // indirect + github.com/stretchr/objx v0.5.2 // indirect + golang.org/x/crypto v0.54.0 // indirect + golang.org/x/net v0.57.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.40.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/third_party/hysteria-core/go.sum b/third_party/hysteria-core/go.sum new file mode 100644 index 0000000..effebfa --- /dev/null +++ b/third_party/hysteria-core/go.sum @@ -0,0 +1,46 @@ +github.com/andybalholm/brotli v1.1.0 h1:eLKJA0d02Lf0mVpIDgYnqXcUn0GqVmEFny3VuID1U3M= +github.com/andybalholm/brotli v1.1.0/go.mod h1:sms7XGricyQI9K10gOSf56VKKWS4oLer58Q+mhRPtnY= +github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= +github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0= +github.com/quic-go/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk= +github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= +github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= +github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEvV+S9iJ2IdQo= +github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= +github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= +github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= +github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= +github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= +go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko= +go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o= +golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= +golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= +golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM= +golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= +golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= +golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= +golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/third_party/hysteria-core/internal/congestion/bbr/bandwidth.go b/third_party/hysteria-core/internal/congestion/bbr/bandwidth.go new file mode 100644 index 0000000..52deb24 --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/bandwidth.go @@ -0,0 +1,27 @@ +package bbr + +import ( + "math" + "time" + + "github.com/apernet/quic-go/congestion" +) + +const ( + infBandwidth = Bandwidth(math.MaxUint64) +) + +// Bandwidth of a connection +type Bandwidth uint64 + +const ( + // BitsPerSecond is 1 bit per second + BitsPerSecond Bandwidth = 1 + // BytesPerSecond is 1 byte per second + BytesPerSecond = 8 * BitsPerSecond +) + +// BandwidthFromDelta calculates the bandwidth from a number of bytes and a time delta +func BandwidthFromDelta(bytes congestion.ByteCount, delta time.Duration) Bandwidth { + return Bandwidth(bytes) * Bandwidth(time.Second) / Bandwidth(delta) * BytesPerSecond +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/bandwidth_sampler.go b/third_party/hysteria-core/internal/congestion/bbr/bandwidth_sampler.go new file mode 100644 index 0000000..2bd66d0 --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/bandwidth_sampler.go @@ -0,0 +1,877 @@ +package bbr + +import ( + "math" + "time" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/monotime" +) + +const ( + infRTT = time.Duration(math.MaxInt64) + defaultConnectionStateMapQueueSize = 256 + defaultCandidatesBufferSize = 256 +) + +type roundTripCount uint64 + +// SendTimeState is a subset of ConnectionStateOnSentPacket which is returned +// to the caller when the packet is acked or lost. +type sendTimeState struct { + // Whether other states in this object is valid. + isValid bool + // Whether the sender is app limited at the time the packet was sent. + // App limited bandwidth sample might be artificially low because the sender + // did not have enough data to send in order to saturate the link. + isAppLimited bool + // Total number of sent bytes at the time the packet was sent. + // Includes the packet itself. + totalBytesSent congestion.ByteCount + // Total number of acked bytes at the time the packet was sent. + totalBytesAcked congestion.ByteCount + // Total number of lost bytes at the time the packet was sent. + totalBytesLost congestion.ByteCount + // Total number of inflight bytes at the time the packet was sent. + // Includes the packet itself. + // It should be equal to |total_bytes_sent| minus the sum of + // |total_bytes_acked|, |total_bytes_lost| and total neutered bytes. + bytesInFlight congestion.ByteCount +} + +func newSendTimeState( + isAppLimited bool, + totalBytesSent congestion.ByteCount, + totalBytesAcked congestion.ByteCount, + totalBytesLost congestion.ByteCount, + bytesInFlight congestion.ByteCount, +) *sendTimeState { + return &sendTimeState{ + isValid: true, + isAppLimited: isAppLimited, + totalBytesSent: totalBytesSent, + totalBytesAcked: totalBytesAcked, + totalBytesLost: totalBytesLost, + bytesInFlight: bytesInFlight, + } +} + +type extraAckedEvent struct { + // The excess bytes acknowlwedged in the time delta for this event. + extraAcked congestion.ByteCount + + // The bytes acknowledged and time delta from the event. + bytesAcked congestion.ByteCount + timeDelta time.Duration + // The round trip of the event. + round roundTripCount +} + +func maxExtraAckedEventFunc(a, b extraAckedEvent) int { + if a.extraAcked > b.extraAcked { + return 1 + } else if a.extraAcked < b.extraAcked { + return -1 + } + return 0 +} + +// BandwidthSample +type bandwidthSample struct { + // The bandwidth at that particular sample. Zero if no valid bandwidth sample + // is available. + bandwidth Bandwidth + // The RTT measurement at this particular sample. Zero if no RTT sample is + // available. Does not correct for delayed ack time. + rtt time.Duration + // |send_rate| is computed from the current packet being acked('P') and an + // earlier packet that is acked before P was sent. + sendRate Bandwidth + // States captured when the packet was sent. + stateAtSend sendTimeState +} + +func newBandwidthSample() *bandwidthSample { + return &bandwidthSample{ + sendRate: infBandwidth, + } +} + +// MaxAckHeightTracker is part of the BandwidthSampler. It is called after every +// ack event to keep track the degree of ack aggregation(a.k.a "ack height"). +type maxAckHeightTracker struct { + // Tracks the maximum number of bytes acked faster than the estimated + // bandwidth. + maxAckHeightFilter *WindowedFilter[extraAckedEvent, roundTripCount] + // The time this aggregation started and the number of bytes acked during it. + aggregationEpochStartTime monotime.Time + aggregationEpochBytes congestion.ByteCount + // The last sent packet number before the current aggregation epoch started. + lastSentPacketNumberBeforeEpoch congestion.PacketNumber + // The number of ack aggregation epochs ever started, including the ongoing + // one. Stats only. + numAckAggregationEpochs uint64 + ackAggregationBandwidthThreshold float64 + startNewAggregationEpochAfterFullRound bool + reduceExtraAckedOnBandwidthIncrease bool +} + +func newMaxAckHeightTracker(windowLength roundTripCount) *maxAckHeightTracker { + return &maxAckHeightTracker{ + maxAckHeightFilter: NewWindowedFilter(windowLength, maxExtraAckedEventFunc), + lastSentPacketNumberBeforeEpoch: invalidPacketNumber, + ackAggregationBandwidthThreshold: 1.0, + } +} + +func (m *maxAckHeightTracker) Get() congestion.ByteCount { + return m.maxAckHeightFilter.GetBest().extraAcked +} + +func (m *maxAckHeightTracker) Update( + bandwidthEstimate Bandwidth, + isNewMaxBandwidth bool, + roundTripCount roundTripCount, + lastSentPacketNumber congestion.PacketNumber, + lastAckedPacketNumber congestion.PacketNumber, + ackTime monotime.Time, + bytesAcked congestion.ByteCount, +) congestion.ByteCount { + forceNewEpoch := false + + if m.reduceExtraAckedOnBandwidthIncrease && isNewMaxBandwidth { + // Save and clear existing entries. + best := m.maxAckHeightFilter.GetBest() + secondBest := m.maxAckHeightFilter.GetSecondBest() + thirdBest := m.maxAckHeightFilter.GetThirdBest() + m.maxAckHeightFilter.Clear() + + // Reinsert the heights into the filter after recalculating. + expectedBytesAcked := bytesFromBandwidthAndTimeDelta(bandwidthEstimate, best.timeDelta) + if expectedBytesAcked < best.bytesAcked { + best.extraAcked = best.bytesAcked - expectedBytesAcked + m.maxAckHeightFilter.Update(best, best.round) + } + expectedBytesAcked = bytesFromBandwidthAndTimeDelta(bandwidthEstimate, secondBest.timeDelta) + if expectedBytesAcked < secondBest.bytesAcked { + secondBest.extraAcked = secondBest.bytesAcked - expectedBytesAcked + m.maxAckHeightFilter.Update(secondBest, secondBest.round) + } + expectedBytesAcked = bytesFromBandwidthAndTimeDelta(bandwidthEstimate, thirdBest.timeDelta) + if expectedBytesAcked < thirdBest.bytesAcked { + thirdBest.extraAcked = thirdBest.bytesAcked - expectedBytesAcked + m.maxAckHeightFilter.Update(thirdBest, thirdBest.round) + } + } + + // If any packet sent after the start of the epoch has been acked, start a new + // epoch. + if m.startNewAggregationEpochAfterFullRound && + m.lastSentPacketNumberBeforeEpoch != invalidPacketNumber && + lastAckedPacketNumber != invalidPacketNumber && + lastAckedPacketNumber > m.lastSentPacketNumberBeforeEpoch { + forceNewEpoch = true + } + if m.aggregationEpochStartTime.IsZero() || forceNewEpoch { + m.aggregationEpochBytes = bytesAcked + m.aggregationEpochStartTime = ackTime + m.lastSentPacketNumberBeforeEpoch = lastSentPacketNumber + m.numAckAggregationEpochs++ + return 0 + } + + // Compute how many bytes are expected to be delivered, assuming max bandwidth + // is correct. + aggregationDelta := ackTime.Sub(m.aggregationEpochStartTime) + expectedBytesAcked := bytesFromBandwidthAndTimeDelta(bandwidthEstimate, aggregationDelta) + // Reset the current aggregation epoch as soon as the ack arrival rate is less + // than or equal to the max bandwidth. + if m.aggregationEpochBytes <= congestion.ByteCount(m.ackAggregationBandwidthThreshold*float64(expectedBytesAcked)) { + // Reset to start measuring a new aggregation epoch. + m.aggregationEpochBytes = bytesAcked + m.aggregationEpochStartTime = ackTime + m.lastSentPacketNumberBeforeEpoch = lastSentPacketNumber + m.numAckAggregationEpochs++ + return 0 + } + + m.aggregationEpochBytes += bytesAcked + + // Compute how many extra bytes were delivered vs max bandwidth. + extraBytesAcked := m.aggregationEpochBytes - expectedBytesAcked + newEvent := extraAckedEvent{ + extraAcked: extraBytesAcked, + bytesAcked: m.aggregationEpochBytes, + timeDelta: aggregationDelta, + } + m.maxAckHeightFilter.Update(newEvent, roundTripCount) + return extraBytesAcked +} + +func (m *maxAckHeightTracker) SetFilterWindowLength(length roundTripCount) { + m.maxAckHeightFilter.SetWindowLength(length) +} + +func (m *maxAckHeightTracker) Reset(newHeight congestion.ByteCount, newTime roundTripCount) { + newEvent := extraAckedEvent{ + extraAcked: newHeight, + round: newTime, + } + m.maxAckHeightFilter.Reset(newEvent, newTime) +} + +func (m *maxAckHeightTracker) SetAckAggregationBandwidthThreshold(threshold float64) { + m.ackAggregationBandwidthThreshold = threshold +} + +func (m *maxAckHeightTracker) SetStartNewAggregationEpochAfterFullRound(value bool) { + m.startNewAggregationEpochAfterFullRound = value +} + +func (m *maxAckHeightTracker) SetReduceExtraAckedOnBandwidthIncrease(value bool) { + m.reduceExtraAckedOnBandwidthIncrease = value +} + +func (m *maxAckHeightTracker) AckAggregationBandwidthThreshold() float64 { + return m.ackAggregationBandwidthThreshold +} + +func (m *maxAckHeightTracker) NumAckAggregationEpochs() uint64 { + return m.numAckAggregationEpochs +} + +// AckPoint represents a point on the ack line. +type ackPoint struct { + ackTime monotime.Time + totalBytesAcked congestion.ByteCount +} + +// RecentAckPoints maintains the most recent 2 ack points at distinct times. +type recentAckPoints struct { + ackPoints [2]ackPoint +} + +func (r *recentAckPoints) Update(ackTime monotime.Time, totalBytesAcked congestion.ByteCount) { + if ackTime.Before(r.ackPoints[1].ackTime) { + r.ackPoints[1].ackTime = ackTime + } else if ackTime.After(r.ackPoints[1].ackTime) { + r.ackPoints[0] = r.ackPoints[1] + r.ackPoints[1].ackTime = ackTime + } + + r.ackPoints[1].totalBytesAcked = totalBytesAcked +} + +func (r *recentAckPoints) Clear() { + r.ackPoints[0] = ackPoint{} + r.ackPoints[1] = ackPoint{} +} + +func (r *recentAckPoints) MostRecentPoint() *ackPoint { + return &r.ackPoints[1] +} + +func (r *recentAckPoints) LessRecentPoint() *ackPoint { + if r.ackPoints[0].totalBytesAcked != 0 { + return &r.ackPoints[0] + } + + return &r.ackPoints[1] +} + +// ConnectionStateOnSentPacket represents the information about a sent packet +// and the state of the connection at the moment the packet was sent, +// specifically the information about the most recently acknowledged packet at +// that moment. +type connectionStateOnSentPacket struct { + // Time at which the packet is sent. + sentTime monotime.Time + // Size of the packet. + size congestion.ByteCount + // The value of |totalBytesSentAtLastAckedPacket| at the time the + // packet was sent. + totalBytesSentAtLastAckedPacket congestion.ByteCount + // The value of |lastAckedPacketSentTime| at the time the packet was + // sent. + lastAckedPacketSentTime monotime.Time + // The value of |lastAckedPacketAckTime| at the time the packet was + // sent. + lastAckedPacketAckTime monotime.Time + // Send time states that are returned to the congestion controller when the + // packet is acked or lost. + sendTimeState sendTimeState +} + +// Snapshot constructor. Records the current state of the bandwidth +// sampler. +// |bytes_in_flight| is the bytes in flight right after the packet is sent. +func newConnectionStateOnSentPacket( + sentTime monotime.Time, + size congestion.ByteCount, + bytesInFlight congestion.ByteCount, + sampler *bandwidthSampler, +) *connectionStateOnSentPacket { + return &connectionStateOnSentPacket{ + sentTime: sentTime, + size: size, + totalBytesSentAtLastAckedPacket: sampler.totalBytesSentAtLastAckedPacket, + lastAckedPacketSentTime: sampler.lastAckedPacketSentTime, + lastAckedPacketAckTime: sampler.lastAckedPacketAckTime, + sendTimeState: *newSendTimeState( + sampler.isAppLimited, + sampler.totalBytesSent, + sampler.totalBytesAcked, + sampler.totalBytesLost, + bytesInFlight, + ), + } +} + +// BandwidthSampler keeps track of sent and acknowledged packets and outputs a +// bandwidth sample for every packet acknowledged. The samples are taken for +// individual packets, and are not filtered; the consumer has to filter the +// bandwidth samples itself. In certain cases, the sampler will locally severely +// underestimate the bandwidth, hence a maximum filter with a size of at least +// one RTT is recommended. +// +// This class bases its samples on the slope of two curves: the number of bytes +// sent over time, and the number of bytes acknowledged as received over time. +// It produces a sample of both slopes for every packet that gets acknowledged, +// based on a slope between two points on each of the corresponding curves. Note +// that due to the packet loss, the number of bytes on each curve might get +// further and further away from each other, meaning that it is not feasible to +// compare byte values coming from different curves with each other. +// +// The obvious points for measuring slope sample are the ones corresponding to +// the packet that was just acknowledged. Let us denote them as S_1 (point at +// which the current packet was sent) and A_1 (point at which the current packet +// was acknowledged). However, taking a slope requires two points on each line, +// so estimating bandwidth requires picking a packet in the past with respect to +// which the slope is measured. +// +// For that purpose, BandwidthSampler always keeps track of the most recently +// acknowledged packet, and records it together with every outgoing packet. +// When a packet gets acknowledged (A_1), it has not only information about when +// it itself was sent (S_1), but also the information about the latest +// acknowledged packet right before it was sent (S_0 and A_0). +// +// Based on that data, send and ack rate are estimated as: +// +// send_rate = (bytes(S_1) - bytes(S_0)) / (time(S_1) - time(S_0)) +// ack_rate = (bytes(A_1) - bytes(A_0)) / (time(A_1) - time(A_0)) +// +// Here, the ack rate is intuitively the rate we want to treat as bandwidth. +// However, in certain cases (e.g. ack compression) the ack rate at a point may +// end up higher than the rate at which the data was originally sent, which is +// not indicative of the real bandwidth. Hence, we use the send rate as an upper +// bound, and the sample value is +// +// rate_sample = min(send_rate, ack_rate) +// +// An important edge case handled by the sampler is tracking the app-limited +// samples. There are multiple meaning of "app-limited" used interchangeably, +// hence it is important to understand and to be able to distinguish between +// them. +// +// Meaning 1: connection state. The connection is said to be app-limited when +// there is no outstanding data to send. This means that certain bandwidth +// samples in the future would not be an accurate indication of the link +// capacity, and it is important to inform consumer about that. Whenever +// connection becomes app-limited, the sampler is notified via OnAppLimited() +// method. +// +// Meaning 2: a phase in the bandwidth sampler. As soon as the bandwidth +// sampler becomes notified about the connection being app-limited, it enters +// app-limited phase. In that phase, all *sent* packets are marked as +// app-limited. Note that the connection itself does not have to be +// app-limited during the app-limited phase, and in fact it will not be +// (otherwise how would it send packets?). The boolean flag below indicates +// whether the sampler is in that phase. +// +// Meaning 3: a flag on the sent packet and on the sample. If a sent packet is +// sent during the app-limited phase, the resulting sample related to the +// packet will be marked as app-limited. +// +// With the terminology issue out of the way, let us consider the question of +// what kind of situation it addresses. +// +// Consider a scenario where we first send packets 1 to 20 at a regular +// bandwidth, and then immediately run out of data. After a few seconds, we send +// packets 21 to 60, and only receive ack for 21 between sending packets 40 and +// 41. In this case, when we sample bandwidth for packets 21 to 40, the S_0/A_0 +// we use to compute the slope is going to be packet 20, a few seconds apart +// from the current packet, hence the resulting estimate would be extremely low +// and not indicative of anything. Only at packet 41 the S_0/A_0 will become 21, +// meaning that the bandwidth sample would exclude the quiescence. +// +// Based on the analysis of that scenario, we implement the following rule: once +// OnAppLimited() is called, all sent packets will produce app-limited samples +// up until an ack for a packet that was sent after OnAppLimited() was called. +// Note that while the scenario above is not the only scenario when the +// connection is app-limited, the approach works in other cases too. + +type congestionEventSample struct { + // The maximum bandwidth sample from all acked packets. + // QuicBandwidth::Zero() if no samples are available. + sampleMaxBandwidth Bandwidth + // Whether |sample_max_bandwidth| is from a app-limited sample. + sampleIsAppLimited bool + // The minimum rtt sample from all acked packets. + // QuicTime::Delta::Infinite() if no samples are available. + sampleRtt time.Duration + // For each packet p in acked packets, this is the max value of INFLIGHT(p), + // where INFLIGHT(p) is the number of bytes acked while p is inflight. + sampleMaxInflight congestion.ByteCount + // The send state of the largest packet in acked_packets, unless it is + // empty. If acked_packets is empty, it's the send state of the largest + // packet in lost_packets. + lastPacketSendState sendTimeState + // The number of extra bytes acked from this ack event, compared to what is + // expected from the flow's bandwidth. Larger value means more ack + // aggregation. + extraAcked congestion.ByteCount +} + +func newCongestionEventSample() *congestionEventSample { + return &congestionEventSample{ + sampleRtt: infRTT, + } +} + +type bandwidthSampler struct { + // The total number of congestion controlled bytes sent during the connection. + totalBytesSent congestion.ByteCount + + // The total number of congestion controlled bytes which were acknowledged. + totalBytesAcked congestion.ByteCount + + // The total number of congestion controlled bytes which were lost. + totalBytesLost congestion.ByteCount + + // The total number of congestion controlled bytes which have been neutered. + totalBytesNeutered congestion.ByteCount + + // The value of |total_bytes_sent_| at the time the last acknowledged packet + // was sent. Valid only when |last_acked_packet_sent_time_| is valid. + totalBytesSentAtLastAckedPacket congestion.ByteCount + + // The time at which the last acknowledged packet was sent. Set to + // QuicTime::Zero() if no valid timestamp is available. + lastAckedPacketSentTime monotime.Time + + // The time at which the most recent packet was acknowledged. + lastAckedPacketAckTime monotime.Time + + // The most recently sent packet. + lastSentPacket congestion.PacketNumber + + // The most recently acked packet. + lastAckedPacket congestion.PacketNumber + + // Indicates whether the bandwidth sampler is currently in an app-limited + // phase. + isAppLimited bool + + // The packet that will be acknowledged after this one will cause the sampler + // to exit the app-limited phase. + endOfAppLimitedPhase congestion.PacketNumber + + // Record of the connection state at the point where each packet in flight was + // sent, indexed by the packet number. + connectionStateMap *packetNumberIndexedQueue[connectionStateOnSentPacket] + + recentAckPoints recentAckPoints + a0Candidates RingBuffer[ackPoint] + + // Maximum number of tracked packets. + maxTrackedPackets congestion.ByteCount + + maxAckHeightTracker *maxAckHeightTracker + totalBytesAckedAfterLastAckEvent congestion.ByteCount + + // True if connection option 'BSAO' is set. + overestimateAvoidance bool + + // True if connection option 'BBRB' is set. + limitMaxAckHeightTrackerBySendRate bool +} + +func newBandwidthSampler(maxAckHeightTrackerWindowLength roundTripCount) *bandwidthSampler { + b := &bandwidthSampler{ + maxAckHeightTracker: newMaxAckHeightTracker(maxAckHeightTrackerWindowLength), + connectionStateMap: newPacketNumberIndexedQueue[connectionStateOnSentPacket](defaultConnectionStateMapQueueSize), + lastSentPacket: invalidPacketNumber, + lastAckedPacket: invalidPacketNumber, + endOfAppLimitedPhase: invalidPacketNumber, + } + + b.a0Candidates.Init(defaultCandidatesBufferSize) + + return b +} + +func (b *bandwidthSampler) MaxAckHeight() congestion.ByteCount { + return b.maxAckHeightTracker.Get() +} + +func (b *bandwidthSampler) NumAckAggregationEpochs() uint64 { + return b.maxAckHeightTracker.NumAckAggregationEpochs() +} + +func (b *bandwidthSampler) SetMaxAckHeightTrackerWindowLength(length roundTripCount) { + b.maxAckHeightTracker.SetFilterWindowLength(length) +} + +func (b *bandwidthSampler) ResetMaxAckHeightTracker(newHeight congestion.ByteCount, newTime roundTripCount) { + b.maxAckHeightTracker.Reset(newHeight, newTime) +} + +func (b *bandwidthSampler) SetStartNewAggregationEpochAfterFullRound(value bool) { + b.maxAckHeightTracker.SetStartNewAggregationEpochAfterFullRound(value) +} + +func (b *bandwidthSampler) SetLimitMaxAckHeightTrackerBySendRate(value bool) { + b.limitMaxAckHeightTrackerBySendRate = value +} + +func (b *bandwidthSampler) SetReduceExtraAckedOnBandwidthIncrease(value bool) { + b.maxAckHeightTracker.SetReduceExtraAckedOnBandwidthIncrease(value) +} + +func (b *bandwidthSampler) EnableOverestimateAvoidance() { + if b.overestimateAvoidance { + return + } + + b.overestimateAvoidance = true + b.maxAckHeightTracker.SetAckAggregationBandwidthThreshold(2.0) +} + +func (b *bandwidthSampler) IsOverestimateAvoidanceEnabled() bool { + return b.overestimateAvoidance +} + +func (b *bandwidthSampler) OnPacketSent( + sentTime monotime.Time, + packetNumber congestion.PacketNumber, + bytes congestion.ByteCount, + bytesInFlight congestion.ByteCount, + isRetransmittable bool, +) { + b.lastSentPacket = packetNumber + + if !isRetransmittable { + return + } + + b.totalBytesSent += bytes + + // If there are no packets in flight, the time at which the new transmission + // opens can be treated as the A_0 point for the purpose of bandwidth + // sampling. This underestimates bandwidth to some extent, and produces some + // artificially low samples for most packets in flight, but it provides with + // samples at important points where we would not have them otherwise, most + // importantly at the beginning of the connection. + if bytesInFlight == 0 { + b.lastAckedPacketAckTime = sentTime + if b.overestimateAvoidance { + b.recentAckPoints.Clear() + b.recentAckPoints.Update(sentTime, b.totalBytesAcked) + b.a0Candidates.Clear() + b.a0Candidates.PushBack(*b.recentAckPoints.MostRecentPoint()) + } + b.totalBytesSentAtLastAckedPacket = b.totalBytesSent + + // In this situation ack compression is not a concern, set send rate to + // effectively infinite. + b.lastAckedPacketSentTime = sentTime + } + + b.connectionStateMap.Emplace(packetNumber, newConnectionStateOnSentPacket( + sentTime, + bytes, + bytesInFlight+bytes, + b, + )) +} + +func (b *bandwidthSampler) OnCongestionEvent( + ackTime monotime.Time, + ackedPackets []congestion.AckedPacketInfo, + lostPackets []congestion.LostPacketInfo, + maxBandwidth Bandwidth, + estBandwidthUpperBound Bandwidth, + roundTripCount roundTripCount, +) congestionEventSample { + eventSample := newCongestionEventSample() + + var lastLostPacketSendState sendTimeState + + for _, p := range lostPackets { + sendState := b.OnPacketLost(p.PacketNumber, p.BytesLost) + if sendState.isValid { + lastLostPacketSendState = sendState + } + } + + if len(ackedPackets) == 0 { + // Only populate send state for a loss-only event. + eventSample.lastPacketSendState = lastLostPacketSendState + return *eventSample + } + + var lastAckedPacketSendState sendTimeState + var maxSendRate Bandwidth + + for _, p := range ackedPackets { + sample := b.onPacketAcknowledged(ackTime, p.PacketNumber) + if !sample.stateAtSend.isValid { + continue + } + + lastAckedPacketSendState = sample.stateAtSend + + if sample.rtt != 0 { + eventSample.sampleRtt = min(eventSample.sampleRtt, sample.rtt) + } + if sample.bandwidth > eventSample.sampleMaxBandwidth { + eventSample.sampleMaxBandwidth = sample.bandwidth + eventSample.sampleIsAppLimited = sample.stateAtSend.isAppLimited + } + if sample.sendRate != infBandwidth { + maxSendRate = max(maxSendRate, sample.sendRate) + } + inflightSample := b.totalBytesAcked - lastAckedPacketSendState.totalBytesAcked + if inflightSample > eventSample.sampleMaxInflight { + eventSample.sampleMaxInflight = inflightSample + } + } + + if !lastLostPacketSendState.isValid { + eventSample.lastPacketSendState = lastAckedPacketSendState + } else if !lastAckedPacketSendState.isValid { + eventSample.lastPacketSendState = lastLostPacketSendState + } else { + // If two packets are inflight and an alarm is armed to lose a packet and it + // wakes up late, then the first of two in flight packets could have been + // acknowledged before the wakeup, which re-evaluates loss detection, and + // could declare the later of the two lost. + if lostPackets[len(lostPackets)-1].PacketNumber > ackedPackets[len(ackedPackets)-1].PacketNumber { + eventSample.lastPacketSendState = lastLostPacketSendState + } else { + eventSample.lastPacketSendState = lastAckedPacketSendState + } + } + + isNewMaxBandwidth := eventSample.sampleMaxBandwidth > maxBandwidth + maxBandwidth = max(maxBandwidth, eventSample.sampleMaxBandwidth) + if b.limitMaxAckHeightTrackerBySendRate { + maxBandwidth = max(maxBandwidth, maxSendRate) + } + + eventSample.extraAcked = b.onAckEventEnd(min(estBandwidthUpperBound, maxBandwidth), isNewMaxBandwidth, roundTripCount) + + return *eventSample +} + +func (b *bandwidthSampler) OnPacketLost(packetNumber congestion.PacketNumber, bytesLost congestion.ByteCount) (s sendTimeState) { + b.totalBytesLost += bytesLost + if sentPacketPointer := b.connectionStateMap.GetEntry(packetNumber); sentPacketPointer != nil { + sentPacketToSendTimeState(sentPacketPointer, &s) + } + return s +} + +func (b *bandwidthSampler) OnPacketNeutered(packetNumber congestion.PacketNumber) { + b.connectionStateMap.Remove(packetNumber, func(sentPacket connectionStateOnSentPacket) { + b.totalBytesNeutered += sentPacket.size + }) +} + +func (b *bandwidthSampler) OnAppLimited() { + b.isAppLimited = true + b.endOfAppLimitedPhase = b.lastSentPacket +} + +func (b *bandwidthSampler) RemoveObsoletePackets(leastUnacked congestion.PacketNumber) { + // A packet can become obsolete when it is removed from QuicUnackedPacketMap's + // view of inflight before it is acked or marked as lost. For example, when + // QuicSentPacketManager::RetransmitCryptoPackets retransmits a crypto packet, + // the packet is removed from QuicUnackedPacketMap's inflight, but is not + // marked as acked or lost in the BandwidthSampler. + b.connectionStateMap.RemoveUpTo(leastUnacked) +} + +func (b *bandwidthSampler) TotalBytesSent() congestion.ByteCount { + return b.totalBytesSent +} + +func (b *bandwidthSampler) TotalBytesLost() congestion.ByteCount { + return b.totalBytesLost +} + +func (b *bandwidthSampler) TotalBytesAcked() congestion.ByteCount { + return b.totalBytesAcked +} + +func (b *bandwidthSampler) TotalBytesNeutered() congestion.ByteCount { + return b.totalBytesNeutered +} + +func (b *bandwidthSampler) IsAppLimited() bool { + return b.isAppLimited +} + +func (b *bandwidthSampler) EndOfAppLimitedPhase() congestion.PacketNumber { + return b.endOfAppLimitedPhase +} + +func (b *bandwidthSampler) max_ack_height() congestion.ByteCount { + return b.maxAckHeightTracker.Get() +} + +func (b *bandwidthSampler) chooseA0Point(totalBytesAcked congestion.ByteCount, a0 *ackPoint) bool { + if b.a0Candidates.Empty() { + return false + } + + if b.a0Candidates.Len() == 1 { + *a0 = *b.a0Candidates.Front() + return true + } + + for i := 1; i < b.a0Candidates.Len(); i++ { + if b.a0Candidates.Offset(i).totalBytesAcked > totalBytesAcked { + *a0 = *b.a0Candidates.Offset(i - 1) + if i > 1 { + for j := 0; j < i-1; j++ { + b.a0Candidates.PopFront() + } + } + return true + } + } + + *a0 = *b.a0Candidates.Back() + for k := 0; k < b.a0Candidates.Len()-1; k++ { + b.a0Candidates.PopFront() + } + return true +} + +func (b *bandwidthSampler) onPacketAcknowledged(ackTime monotime.Time, packetNumber congestion.PacketNumber) bandwidthSample { + sample := newBandwidthSample() + b.lastAckedPacket = packetNumber + sentPacketPointer := b.connectionStateMap.GetEntry(packetNumber) + if sentPacketPointer == nil { + return *sample + } + + // OnPacketAcknowledgedInner + b.totalBytesAcked += sentPacketPointer.size + b.totalBytesSentAtLastAckedPacket = sentPacketPointer.sendTimeState.totalBytesSent + b.lastAckedPacketSentTime = sentPacketPointer.sentTime + b.lastAckedPacketAckTime = ackTime + if b.overestimateAvoidance { + b.recentAckPoints.Update(ackTime, b.totalBytesAcked) + } + + if b.isAppLimited { + // Exit app-limited phase in two cases: + // (1) end_of_app_limited_phase_ is not initialized, i.e., so far all + // packets are sent while there are buffered packets or pending data. + // (2) The current acked packet is after the sent packet marked as the end + // of the app limit phase. + if b.endOfAppLimitedPhase == invalidPacketNumber || + packetNumber > b.endOfAppLimitedPhase { + b.isAppLimited = false + } + } + + // There might have been no packets acknowledged at the moment when the + // current packet was sent. In that case, there is no bandwidth sample to + // make. + if sentPacketPointer.lastAckedPacketSentTime.IsZero() { + return *sample + } + + // Infinite rate indicates that the sampler is supposed to discard the + // current send rate sample and use only the ack rate. + sendRate := infBandwidth + if sentPacketPointer.sentTime.After(sentPacketPointer.lastAckedPacketSentTime) { + sendRate = BandwidthFromDelta( + sentPacketPointer.sendTimeState.totalBytesSent-sentPacketPointer.totalBytesSentAtLastAckedPacket, + sentPacketPointer.sentTime.Sub(sentPacketPointer.lastAckedPacketSentTime), + ) + } + + var a0 ackPoint + if b.overestimateAvoidance && b.chooseA0Point(sentPacketPointer.sendTimeState.totalBytesAcked, &a0) { + } else { + a0.ackTime = sentPacketPointer.lastAckedPacketAckTime + a0.totalBytesAcked = sentPacketPointer.sendTimeState.totalBytesAcked + } + + // During the slope calculation, ensure that ack time of the current packet is + // always larger than the time of the previous packet, otherwise division by + // zero or integer underflow can occur. + if ackTime.Sub(a0.ackTime) <= 0 { + return *sample + } + + ackRate := BandwidthFromDelta(b.totalBytesAcked-a0.totalBytesAcked, ackTime.Sub(a0.ackTime)) + + sample.bandwidth = min(sendRate, ackRate) + // Note: this sample does not account for delayed acknowledgement time. This + // means that the RTT measurements here can be artificially high, especially + // on low bandwidth connections. + sample.rtt = ackTime.Sub(sentPacketPointer.sentTime) + sample.sendRate = sendRate + sentPacketToSendTimeState(sentPacketPointer, &sample.stateAtSend) + + return *sample +} + +func (b *bandwidthSampler) onAckEventEnd( + bandwidthEstimate Bandwidth, + isNewMaxBandwidth bool, + roundTripCount roundTripCount, +) congestion.ByteCount { + newlyAckedBytes := b.totalBytesAcked - b.totalBytesAckedAfterLastAckEvent + if newlyAckedBytes == 0 { + return 0 + } + b.totalBytesAckedAfterLastAckEvent = b.totalBytesAcked + extraAcked := b.maxAckHeightTracker.Update( + bandwidthEstimate, + isNewMaxBandwidth, + roundTripCount, + b.lastSentPacket, + b.lastAckedPacket, + b.lastAckedPacketAckTime, + newlyAckedBytes, + ) + // If |extra_acked| is zero, i.e. this ack event marks the start of a new ack + // aggregation epoch, save LessRecentPoint, which is the last ack point of the + // previous epoch, as a A0 candidate. + if b.overestimateAvoidance && extraAcked == 0 { + b.a0Candidates.PushBack(*b.recentAckPoints.LessRecentPoint()) + } + return extraAcked +} + +func sentPacketToSendTimeState(sentPacket *connectionStateOnSentPacket, sendTimeState *sendTimeState) { + *sendTimeState = sentPacket.sendTimeState + sendTimeState.isValid = true +} + +// BytesFromBandwidthAndTimeDelta calculates the bytes +// from a bandwidth(bits per second) and a time delta +func bytesFromBandwidthAndTimeDelta(bandwidth Bandwidth, delta time.Duration) congestion.ByteCount { + return (congestion.ByteCount(bandwidth) * congestion.ByteCount(delta)) / + (congestion.ByteCount(time.Second) * 8) +} + +func timeDeltaFromBytesAndBandwidth(bytes congestion.ByteCount, bandwidth Bandwidth) time.Duration { + return time.Duration(bytes*8) * time.Second / time.Duration(bandwidth) +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/bbr_sender.go b/third_party/hysteria-core/internal/congestion/bbr/bbr_sender.go new file mode 100644 index 0000000..4185311 --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/bbr_sender.go @@ -0,0 +1,1088 @@ +package bbr + +import ( + "fmt" + "math/rand" + "net" + "os" + "strconv" + "strings" + "time" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/monotime" + + "github.com/apernet/hysteria/core/v2/internal/congestion/common" +) + +// BbrSender implements BBR congestion control algorithm. BBR aims to estimate +// the current available Bottleneck Bandwidth and RTT (hence the name), and +// regulates the pacing rate and the size of the congestion window based on +// those signals. +// +// BBR relies on pacing in order to function properly. Do not use BBR when +// pacing is disabled. +// + +const ( + minBps = 65536 // 64 KB/s + + invalidPacketNumber = -1 + initialCongestionWindowPackets = 32 + minCongestionWindowPackets = 4 + + // Constants based on TCP defaults. + // The minimum CWND to ensure delayed acks don't reduce bandwidth measurements. + // Does not inflate the pacing rate. + // The gain used for the STARTUP, equal to 2/ln(2). + defaultHighGain = 2.885 + // The newly derived CWND gain for STARTUP, 2. + derivedHighCWNDGain = 2.0 + + debugEnv = "HYSTERIA_BBR_DEBUG" +) + +// The cycle of gains used during the PROBE_BW stage. +var pacingGain = [...]float64{1.25, 0.75, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0} + +const ( + // The length of the gain cycle. + gainCycleLength = len(pacingGain) + // The size of the bandwidth filter window, in round-trips. + bandwidthWindowSize = gainCycleLength + 2 + + // The time after which the current min_rtt value expires. + minRttExpiry = 10 * time.Second + // The minimum time the connection can spend in PROBE_RTT mode. + probeRttTime = 200 * time.Millisecond + // If the bandwidth does not increase by the factor of |kStartupGrowthTarget| + // within |kRoundTripsWithoutGrowthBeforeExitingStartup| rounds, the connection + // will exit the STARTUP mode. + startupGrowthTarget = 1.25 + roundTripsWithoutGrowthBeforeExitingStartup = int64(3) + + // Flag. + defaultStartupFullLossCount = 8 + quicBbr2DefaultLossThreshold = 0.02 +) + +type bbrMode int + +const ( + // Startup phase of the connection. + bbrModeStartup = iota + // After achieving the highest possible bandwidth during the startup, lower + // the pacing rate in order to drain the queue. + bbrModeDrain + // Cruising mode. + bbrModeProbeBw + // Temporarily slow down sending in order to empty the buffer and measure + // the real minimum RTT. + bbrModeProbeRtt +) + +// Indicates how the congestion control limits the amount of bytes in flight. +type bbrRecoveryState int + +const ( + // Do not limit. + bbrRecoveryStateNotInRecovery = iota + // Allow an extra outstanding byte for each byte acknowledged. + bbrRecoveryStateConservation + // Allow two extra outstanding bytes for each byte acknowledged (slow + // start). + bbrRecoveryStateGrowth +) + +type Profile string + +const ( + ProfileConservative Profile = "conservative" + ProfileStandard Profile = "standard" + ProfileAggressive Profile = "aggressive" +) + +type profileConfig struct { + highGain float64 + highCwndGain float64 + congestionWindowGainConstant float64 + numStartupRtts int64 + drainToTarget bool + detectOvershooting bool + bytesLostMultiplier uint8 + enableAckAggregationStartup bool + expireAckAggregationStartup bool + enableOverestimateAvoidance bool + reduceExtraAckedOnBandwidthIncrease bool +} + +func ParseProfile(profile string) (Profile, error) { + switch normalized := strings.ToLower(profile); normalized { + case "", string(ProfileStandard): + return ProfileStandard, nil + case string(ProfileConservative): + return ProfileConservative, nil + case string(ProfileAggressive): + return ProfileAggressive, nil + default: + return "", fmt.Errorf("unsupported BBR profile %q", profile) + } +} + +func configForProfile(profile Profile) profileConfig { + switch profile { + case ProfileConservative: + return profileConfig{ + highGain: 2.25, + highCwndGain: 1.75, + congestionWindowGainConstant: 1.75, + numStartupRtts: 2, + drainToTarget: true, + detectOvershooting: true, + bytesLostMultiplier: 1, + enableOverestimateAvoidance: true, + reduceExtraAckedOnBandwidthIncrease: true, + } + case ProfileAggressive: + return profileConfig{ + highGain: 3.0, + highCwndGain: 2.25, + congestionWindowGainConstant: 2.5, + numStartupRtts: 4, + bytesLostMultiplier: 2, + enableAckAggregationStartup: true, + expireAckAggregationStartup: true, + } + default: + return profileConfig{ + highGain: defaultHighGain, + highCwndGain: derivedHighCWNDGain, + congestionWindowGainConstant: 2.0, + numStartupRtts: roundTripsWithoutGrowthBeforeExitingStartup, + bytesLostMultiplier: 2, + } + } +} + +type bbrSender struct { + rttStats congestion.RTTStatsProvider + clock Clock + pacer *common.Pacer + + mode bbrMode + + // Bandwidth sampler provides BBR with the bandwidth measurements at + // individual points. + sampler *bandwidthSampler + + // The number of the round trips that have occurred during the connection. + roundTripCount roundTripCount + + // The packet number of the most recently sent packet. + lastSentPacket congestion.PacketNumber + // Acknowledgement of any packet after |current_round_trip_end_| will cause + // the round trip counter to advance. + currentRoundTripEnd congestion.PacketNumber + + // Number of congestion events with some losses, in the current round. + numLossEventsInRound uint64 + + // Number of total bytes lost in the current round. + bytesLostInRound congestion.ByteCount + + // The filter that tracks the maximum bandwidth over the multiple recent + // round-trips. + maxBandwidth *WindowedFilter[Bandwidth, roundTripCount] + + // Minimum RTT estimate. Automatically expires within 10 seconds (and + // triggers PROBE_RTT mode) if no new value is sampled during that period. + minRtt time.Duration + // The time at which the current value of |min_rtt_| was assigned. + minRttTimestamp monotime.Time + + // The maximum allowed number of bytes in flight. + congestionWindow congestion.ByteCount + + // The initial value of the |congestion_window_|. + initialCongestionWindow congestion.ByteCount + + // The largest value the |congestion_window_| can achieve. + maxCongestionWindow congestion.ByteCount + + // The smallest value the |congestion_window_| can achieve. + minCongestionWindow congestion.ByteCount + + // The BBR profile used by the sender. + profile Profile + + // The pacing gain applied during the STARTUP phase. + highGain float64 + + // The CWND gain applied during the STARTUP phase. + highCwndGain float64 + + // The pacing gain applied during the DRAIN phase. + drainGain float64 + + // The current pacing rate of the connection. + pacingRate Bandwidth + + // The gain currently applied to the pacing rate. + pacingGain float64 + // The gain currently applied to the congestion window. + congestionWindowGain float64 + + // The gain used for the congestion window during PROBE_BW. Latched from + // quic_bbr_cwnd_gain flag. + congestionWindowGainConstant float64 + // The number of RTTs to stay in STARTUP mode. Defaults to 3. + numStartupRtts int64 + + // Number of round-trips in PROBE_BW mode, used for determining the current + // pacing gain cycle. + cycleCurrentOffset int + // The time at which the last pacing gain cycle was started. + lastCycleStart monotime.Time + + // Indicates whether the connection has reached the full bandwidth mode. + isAtFullBandwidth bool + // Number of rounds during which there was no significant bandwidth increase. + roundsWithoutBandwidthGain int64 + // The bandwidth compared to which the increase is measured. + bandwidthAtLastRound Bandwidth + + // Set to true upon exiting quiescence. + exitingQuiescence bool + + // Time at which PROBE_RTT has to be exited. Setting it to zero indicates + // that the time is yet unknown as the number of packets in flight has not + // reached the required value. + exitProbeRttAt monotime.Time + // Indicates whether a round-trip has passed since PROBE_RTT became active. + probeRttRoundPassed bool + + // Indicates whether the most recent bandwidth sample was marked as + // app-limited. + lastSampleIsAppLimited bool + // Indicates whether any non app-limited samples have been recorded. + hasNoAppLimitedSample bool + + // Current state of recovery. + recoveryState bbrRecoveryState + // Receiving acknowledgement of a packet after |end_recovery_at_| will cause + // BBR to exit the recovery mode. A value above zero indicates at least one + // loss has been detected, so it must not be set back to zero. + endRecoveryAt congestion.PacketNumber + // A window used to limit the number of bytes in flight during loss recovery. + recoveryWindow congestion.ByteCount + // If true, consider all samples in recovery app-limited. + isAppLimitedRecovery bool // not used + + // When true, pace at 1.5x and disable packet conservation in STARTUP. + slowerStartup bool // not used + // When true, disables packet conservation in STARTUP. + rateBasedStartup bool // not used + + // When true, add the most recent ack aggregation measurement during STARTUP. + enableAckAggregationDuringStartup bool + // When true, expire the windowed ack aggregation values in STARTUP when + // bandwidth increases more than 25%. + expireAckAggregationInStartup bool + + // If true, will not exit low gain mode until bytes_in_flight drops below BDP + // or it's time for high gain mode. + drainToTarget bool + + // If true, slow down pacing rate in STARTUP when overshooting is detected. + detectOvershooting bool + // Bytes lost while detect_overshooting_ is true. + bytesLostWhileDetectingOvershooting congestion.ByteCount + // Slow down pacing rate if + // bytes_lost_while_detecting_overshooting_ * + // bytes_lost_multiplier_while_detecting_overshooting_ > IW. + bytesLostMultiplierWhileDetectingOvershooting uint8 + // When overshooting is detected, do not drop pacing_rate_ below this value / + // min_rtt. + cwndToCalculateMinPacingRate congestion.ByteCount + + // Max congestion window when adjusting network parameters. + maxCongestionWindowWithNetworkParametersAdjusted congestion.ByteCount // not used + + // Params. + maxDatagramSize congestion.ByteCount + // Recorded on packet sent. equivalent |unacked_packets_->bytes_in_flight()| + bytesInFlight congestion.ByteCount + + debug bool +} + +var _ congestion.CongestionControl = &bbrSender{} + +func NewBbrSender( + clock Clock, + initialMaxDatagramSize congestion.ByteCount, + profile Profile, +) *bbrSender { + return newBbrSender( + clock, + initialMaxDatagramSize, + initialCongestionWindowPackets*initialMaxDatagramSize, + congestion.MaxCongestionWindowPackets*initialMaxDatagramSize, + profile, + ) +} + +func newBbrSender( + clock Clock, + initialMaxDatagramSize, + initialCongestionWindow, + initialMaxCongestionWindow congestion.ByteCount, + profile Profile, +) *bbrSender { + debug, _ := strconv.ParseBool(os.Getenv(debugEnv)) + b := &bbrSender{ + clock: clock, + mode: bbrModeStartup, + sampler: newBandwidthSampler(roundTripCount(bandwidthWindowSize)), + lastSentPacket: invalidPacketNumber, + currentRoundTripEnd: invalidPacketNumber, + maxBandwidth: NewWindowedFilter(roundTripCount(bandwidthWindowSize), MaxFilter[Bandwidth]), + congestionWindow: initialCongestionWindow, + initialCongestionWindow: initialCongestionWindow, + maxCongestionWindow: initialMaxCongestionWindow, + minCongestionWindow: minCongestionWindowForMaxDatagramSize(initialMaxDatagramSize), + profile: ProfileStandard, + highGain: defaultHighGain, + highCwndGain: derivedHighCWNDGain, + drainGain: 1.0 / defaultHighGain, + pacingGain: 1.0, + congestionWindowGain: 1.0, + congestionWindowGainConstant: 2.0, + numStartupRtts: roundTripsWithoutGrowthBeforeExitingStartup, + recoveryState: bbrRecoveryStateNotInRecovery, + endRecoveryAt: invalidPacketNumber, + recoveryWindow: initialMaxCongestionWindow, + bytesLostMultiplierWhileDetectingOvershooting: 2, + cwndToCalculateMinPacingRate: initialCongestionWindow, + maxCongestionWindowWithNetworkParametersAdjusted: initialMaxCongestionWindow, + maxDatagramSize: initialMaxDatagramSize, + debug: debug, + } + b.pacer = common.NewPacer(b.bandwidthForPacer) + b.applyProfile(profile) + if b.debug { + b.debugPrint("Profile: %s", b.profile) + } + + b.enterStartupMode(b.clock.Now()) + + return b +} + +func (b *bbrSender) applyProfile(profile Profile) { + if profile == "" { + profile = ProfileStandard + } + cfg := configForProfile(profile) + b.profile = profile + b.highGain = cfg.highGain + b.highCwndGain = cfg.highCwndGain + b.drainGain = 1.0 / cfg.highGain + b.congestionWindowGainConstant = cfg.congestionWindowGainConstant + b.numStartupRtts = cfg.numStartupRtts + b.drainToTarget = cfg.drainToTarget + b.detectOvershooting = cfg.detectOvershooting + b.bytesLostMultiplierWhileDetectingOvershooting = cfg.bytesLostMultiplier + b.enableAckAggregationDuringStartup = cfg.enableAckAggregationStartup + b.expireAckAggregationInStartup = cfg.expireAckAggregationStartup + if cfg.enableOverestimateAvoidance { + b.sampler.EnableOverestimateAvoidance() + } + b.sampler.SetReduceExtraAckedOnBandwidthIncrease(cfg.reduceExtraAckedOnBandwidthIncrease) +} + +func minCongestionWindowForMaxDatagramSize(maxDatagramSize congestion.ByteCount) congestion.ByteCount { + return minCongestionWindowPackets * maxDatagramSize +} + +func scaleByteWindowForDatagramSize(window, oldMaxDatagramSize, newMaxDatagramSize congestion.ByteCount) congestion.ByteCount { + if oldMaxDatagramSize == newMaxDatagramSize { + return window + } + return congestion.ByteCount(uint64(window) * uint64(newMaxDatagramSize) / uint64(oldMaxDatagramSize)) +} + +func (b *bbrSender) rescalePacketSizedWindows(maxDatagramSize congestion.ByteCount) { + oldMaxDatagramSize := b.maxDatagramSize + b.maxDatagramSize = maxDatagramSize + b.initialCongestionWindow = scaleByteWindowForDatagramSize(b.initialCongestionWindow, oldMaxDatagramSize, maxDatagramSize) + b.maxCongestionWindow = scaleByteWindowForDatagramSize(b.maxCongestionWindow, oldMaxDatagramSize, maxDatagramSize) + b.minCongestionWindow = minCongestionWindowForMaxDatagramSize(maxDatagramSize) + b.cwndToCalculateMinPacingRate = scaleByteWindowForDatagramSize(b.cwndToCalculateMinPacingRate, oldMaxDatagramSize, maxDatagramSize) + b.maxCongestionWindowWithNetworkParametersAdjusted = scaleByteWindowForDatagramSize( + b.maxCongestionWindowWithNetworkParametersAdjusted, + oldMaxDatagramSize, + maxDatagramSize, + ) +} + +func (b *bbrSender) SetRTTStatsProvider(provider congestion.RTTStatsProvider) { + b.rttStats = provider +} + +// TimeUntilSend implements the SendAlgorithm interface. +func (b *bbrSender) TimeUntilSend(bytesInFlight congestion.ByteCount) monotime.Time { + return b.pacer.TimeUntilSend() +} + +// HasPacingBudget implements the SendAlgorithm interface. +func (b *bbrSender) HasPacingBudget(now monotime.Time) bool { + return b.pacer.Budget(now) >= b.maxDatagramSize +} + +// OnPacketSent implements the SendAlgorithm interface. +func (b *bbrSender) OnPacketSent( + sentTime monotime.Time, + bytesInFlight congestion.ByteCount, + packetNumber congestion.PacketNumber, + bytes congestion.ByteCount, + isRetransmittable bool, +) { + b.pacer.SentPacket(sentTime, bytes) + + b.lastSentPacket = packetNumber + b.bytesInFlight = bytesInFlight + + if bytesInFlight == 0 { + b.exitingQuiescence = true + } + + b.sampler.OnPacketSent(sentTime, packetNumber, bytes, bytesInFlight, isRetransmittable) +} + +// CanSend implements the SendAlgorithm interface. +func (b *bbrSender) CanSend(bytesInFlight congestion.ByteCount) bool { + return bytesInFlight < b.GetCongestionWindow() +} + +// MaybeExitSlowStart implements the SendAlgorithm interface. +func (b *bbrSender) MaybeExitSlowStart() { + // Do nothing +} + +// OnPacketAcked implements the SendAlgorithm interface. +func (b *bbrSender) OnPacketAcked(number congestion.PacketNumber, ackedBytes, priorInFlight congestion.ByteCount, eventTime monotime.Time) { + // Do nothing. +} + +// OnPacketLost implements the SendAlgorithm interface. +func (b *bbrSender) OnPacketLost(number congestion.PacketNumber, lostBytes, priorInFlight congestion.ByteCount) { + // Do nothing. +} + +// OnRetransmissionTimeout implements the SendAlgorithm interface. +func (b *bbrSender) OnRetransmissionTimeout(packetsRetransmitted bool) { + // Do nothing. +} + +// SetMaxDatagramSize implements the SendAlgorithm interface. +func (b *bbrSender) SetMaxDatagramSize(s congestion.ByteCount) { + if b.debug { + b.debugPrint("Max Datagram Size: %d", s) + } + if s < b.maxDatagramSize { + panic(fmt.Sprintf("congestion BUG: decreased max datagram size from %d to %d", b.maxDatagramSize, s)) + } + oldMinCongestionWindow := b.minCongestionWindow + oldInitialCongestionWindow := b.initialCongestionWindow + b.rescalePacketSizedWindows(s) + switch b.congestionWindow { + case oldMinCongestionWindow: + b.congestionWindow = b.minCongestionWindow + case oldInitialCongestionWindow: + b.congestionWindow = b.initialCongestionWindow + default: + b.congestionWindow = min(b.maxCongestionWindow, max(b.congestionWindow, b.minCongestionWindow)) + } + b.recoveryWindow = min(b.maxCongestionWindow, max(b.recoveryWindow, b.minCongestionWindow)) + b.pacer.SetMaxDatagramSize(s) +} + +// InSlowStart implements the SendAlgorithmWithDebugInfos interface. +func (b *bbrSender) InSlowStart() bool { + return b.mode == bbrModeStartup +} + +// InRecovery implements the SendAlgorithmWithDebugInfos interface. +func (b *bbrSender) InRecovery() bool { + return b.recoveryState != bbrRecoveryStateNotInRecovery +} + +// GetCongestionWindow implements the SendAlgorithmWithDebugInfos interface. +func (b *bbrSender) GetCongestionWindow() congestion.ByteCount { + if b.mode == bbrModeProbeRtt { + return b.probeRttCongestionWindow() + } + + if b.InRecovery() { + return min(b.congestionWindow, b.recoveryWindow) + } + + return b.congestionWindow +} + +func (b *bbrSender) OnCongestionEvent(number congestion.PacketNumber, lostBytes, priorInFlight congestion.ByteCount) { + // Do nothing. +} + +func (b *bbrSender) OnCongestionEventEx(priorInFlight congestion.ByteCount, eventTime monotime.Time, ackedPackets []congestion.AckedPacketInfo, lostPackets []congestion.LostPacketInfo) { + totalBytesAckedBefore := b.sampler.TotalBytesAcked() + totalBytesLostBefore := b.sampler.TotalBytesLost() + + var isRoundStart, minRttExpired bool + var excessAcked, bytesLost congestion.ByteCount + + // The send state of the largest packet in acked_packets, unless it is + // empty. If acked_packets is empty, it's the send state of the largest + // packet in lost_packets. + var lastPacketSendState sendTimeState + + b.maybeAppLimited(priorInFlight) + + // Update bytesInFlight + b.bytesInFlight = priorInFlight + for _, p := range ackedPackets { + b.bytesInFlight -= p.BytesAcked + } + for _, p := range lostPackets { + b.bytesInFlight -= p.BytesLost + } + + if len(ackedPackets) != 0 { + lastAckedPacket := ackedPackets[len(ackedPackets)-1].PacketNumber + isRoundStart = b.updateRoundTripCounter(lastAckedPacket) + b.updateRecoveryState(lastAckedPacket, len(lostPackets) != 0, isRoundStart) + } + + sample := b.sampler.OnCongestionEvent(eventTime, + ackedPackets, lostPackets, b.maxBandwidth.GetBest(), infBandwidth, b.roundTripCount) + if sample.lastPacketSendState.isValid { + b.lastSampleIsAppLimited = sample.lastPacketSendState.isAppLimited + b.hasNoAppLimitedSample = b.hasNoAppLimitedSample || !b.lastSampleIsAppLimited + } + // Avoid updating |max_bandwidth_| if a) this is a loss-only event, or b) all + // packets in |acked_packets| did not generate valid samples. (e.g. ack of + // ack-only packets). In both cases, sampler_.total_bytes_acked() will not + // change. + if totalBytesAckedBefore != b.sampler.TotalBytesAcked() { + if !sample.sampleIsAppLimited || sample.sampleMaxBandwidth > b.maxBandwidth.GetBest() { + b.maxBandwidth.Update(sample.sampleMaxBandwidth, b.roundTripCount) + } + } + + if sample.sampleRtt != infRTT { + minRttExpired = b.maybeUpdateMinRtt(eventTime, sample.sampleRtt) + } + bytesLost = b.sampler.TotalBytesLost() - totalBytesLostBefore + + excessAcked = sample.extraAcked + lastPacketSendState = sample.lastPacketSendState + + if len(lostPackets) != 0 { + b.numLossEventsInRound++ + b.bytesLostInRound += bytesLost + } + + // Handle logic specific to PROBE_BW mode. + if b.mode == bbrModeProbeBw { + b.updateGainCyclePhase(eventTime, priorInFlight, len(lostPackets) != 0) + } + + // Handle logic specific to STARTUP and DRAIN modes. + if isRoundStart && !b.isAtFullBandwidth { + b.checkIfFullBandwidthReached(&lastPacketSendState) + } + + b.maybeExitStartupOrDrain(eventTime) + + // Handle logic specific to PROBE_RTT. + b.maybeEnterOrExitProbeRtt(eventTime, isRoundStart, minRttExpired) + + // Calculate number of packets acked and lost. + bytesAcked := b.sampler.TotalBytesAcked() - totalBytesAckedBefore + + // After the model is updated, recalculate the pacing rate and congestion + // window. + b.calculatePacingRate(bytesLost) + b.calculateCongestionWindow(bytesAcked, excessAcked) + b.calculateRecoveryWindow(bytesAcked, bytesLost) + + // Cleanup internal state. + // This is where we clean up obsolete (acked or lost) packets from the bandwidth sampler. + // The "least unacked" should actually be FirstOutstanding, but since we are not passing + // that through OnCongestionEventEx, we will only do an estimate using acked/lost packets + // for now. Because of fast retransmission, they should differ by no more than 2 packets. + // (this is controlled by packetThreshold in quic-go's sentPacketHandler) + var leastUnacked congestion.PacketNumber + if len(ackedPackets) != 0 { + leastUnacked = ackedPackets[len(ackedPackets)-1].PacketNumber - 2 + } else { + leastUnacked = lostPackets[len(lostPackets)-1].PacketNumber + 1 + } + b.sampler.RemoveObsoletePackets(leastUnacked) + + if isRoundStart { + b.numLossEventsInRound = 0 + b.bytesLostInRound = 0 + } +} + +func (b *bbrSender) PacingRate() Bandwidth { + if b.pacingRate == 0 { + return Bandwidth(b.highGain * float64( + BandwidthFromDelta(b.initialCongestionWindow, b.getMinRtt()), + )) + } + + return b.pacingRate +} + +// Sets the CWND gain used in STARTUP. Must be greater than 1. +func (b *bbrSender) setHighCwndGain(highCwndGain float64) { + b.highCwndGain = highCwndGain + if b.mode == bbrModeStartup { + b.congestionWindowGain = highCwndGain + } +} + +// Get the current bandwidth estimate. Note that Bandwidth is in bits per second. +func (b *bbrSender) bandwidthEstimate() Bandwidth { + return b.maxBandwidth.GetBest() +} + +func (b *bbrSender) bandwidthForPacer() congestion.ByteCount { + bps := congestion.ByteCount(float64(b.PacingRate()) / float64(BytesPerSecond)) + if bps < minBps { + // We need to make sure that the bandwidth value for pacer is never zero, + // otherwise it will go into an edge case where HasPacingBudget = false + // but TimeUntilSend is before, causing the quic-go send loop to go crazy and get stuck. + return minBps + } + return bps +} + +// Returns the current estimate of the RTT of the connection. Outside of the +// edge cases, this is minimum RTT. +func (b *bbrSender) getMinRtt() time.Duration { + if b.minRtt != 0 { + return b.minRtt + } + // min_rtt could be available if the handshake packet gets neutered then + // gets acknowledged. This could only happen for QUIC crypto where we do not + // drop keys. + minRtt := b.rttStats.MinRTT() + if minRtt == 0 { + return 100 * time.Millisecond + } else { + return minRtt + } +} + +// Computes the target congestion window using the specified gain. +func (b *bbrSender) getTargetCongestionWindow(gain float64) congestion.ByteCount { + bdp := bdpFromRttAndBandwidth(b.getMinRtt(), b.bandwidthEstimate()) + congestionWindow := congestion.ByteCount(gain * float64(bdp)) + + // BDP estimate will be zero if no bandwidth samples are available yet. + if congestionWindow == 0 { + congestionWindow = congestion.ByteCount(gain * float64(b.initialCongestionWindow)) + } + + return max(congestionWindow, b.minCongestionWindow) +} + +// The target congestion window during PROBE_RTT. +func (b *bbrSender) probeRttCongestionWindow() congestion.ByteCount { + return b.minCongestionWindow +} + +func (b *bbrSender) maybeUpdateMinRtt(now monotime.Time, sampleMinRtt time.Duration) bool { + // Do not expire min_rtt if none was ever available. + minRttExpired := b.minRtt != 0 && now.After(b.minRttTimestamp.Add(minRttExpiry)) + if minRttExpired || sampleMinRtt < b.minRtt || b.minRtt == 0 { + b.minRtt = sampleMinRtt + b.minRttTimestamp = now + } + + return minRttExpired +} + +// Enters the STARTUP mode. +func (b *bbrSender) enterStartupMode(now monotime.Time) { + b.mode = bbrModeStartup + // b.maybeTraceStateChange(logging.CongestionStateStartup) + b.pacingGain = b.highGain + b.congestionWindowGain = b.highCwndGain + + if b.debug { + b.debugPrint("Phase: STARTUP") + } +} + +// Enters the PROBE_BW mode. +func (b *bbrSender) enterProbeBandwidthMode(now monotime.Time) { + b.mode = bbrModeProbeBw + // b.maybeTraceStateChange(logging.CongestionStateProbeBw) + b.congestionWindowGain = b.congestionWindowGainConstant + + // Pick a random offset for the gain cycle out of {0, 2..7} range. 1 is + // excluded because in that case increased gain and decreased gain would not + // follow each other. + b.cycleCurrentOffset = int(rand.Int31n(congestion.PacketsPerConnectionID)) % (gainCycleLength - 1) + if b.cycleCurrentOffset >= 1 { + b.cycleCurrentOffset += 1 + } + + b.lastCycleStart = now + b.pacingGain = pacingGain[b.cycleCurrentOffset] + + if b.debug { + b.debugPrint("Phase: PROBE_BW") + } +} + +// Updates the round-trip counter if a round-trip has passed. Returns true if +// the counter has been advanced. +func (b *bbrSender) updateRoundTripCounter(lastAckedPacket congestion.PacketNumber) bool { + if b.currentRoundTripEnd == invalidPacketNumber || lastAckedPacket > b.currentRoundTripEnd { + b.roundTripCount++ + b.currentRoundTripEnd = b.lastSentPacket + return true + } + return false +} + +// Updates the current gain used in PROBE_BW mode. +func (b *bbrSender) updateGainCyclePhase(now monotime.Time, priorInFlight congestion.ByteCount, hasLosses bool) { + // In most cases, the cycle is advanced after an RTT passes. + shouldAdvanceGainCycling := now.After(b.lastCycleStart.Add(b.getMinRtt())) + // If the pacing gain is above 1.0, the connection is trying to probe the + // bandwidth by increasing the number of bytes in flight to at least + // pacing_gain * BDP. Make sure that it actually reaches the target, as long + // as there are no losses suggesting that the buffers are not able to hold + // that much. + if b.pacingGain > 1.0 && !hasLosses && priorInFlight < b.getTargetCongestionWindow(b.pacingGain) { + shouldAdvanceGainCycling = false + } + + // If pacing gain is below 1.0, the connection is trying to drain the extra + // queue which could have been incurred by probing prior to it. If the number + // of bytes in flight falls down to the estimated BDP value earlier, conclude + // that the queue has been successfully drained and exit this cycle early. + if b.pacingGain < 1.0 && b.bytesInFlight <= b.getTargetCongestionWindow(1) { + shouldAdvanceGainCycling = true + } + + if shouldAdvanceGainCycling { + b.cycleCurrentOffset = (b.cycleCurrentOffset + 1) % gainCycleLength + b.lastCycleStart = now + // Stay in low gain mode until the target BDP is hit. + // Low gain mode will be exited immediately when the target BDP is achieved. + if b.drainToTarget && b.pacingGain < 1 && + pacingGain[b.cycleCurrentOffset] == 1 && + b.bytesInFlight > b.getTargetCongestionWindow(1) { + return + } + b.pacingGain = pacingGain[b.cycleCurrentOffset] + } +} + +// Tracks for how many round-trips the bandwidth has not increased +// significantly. +func (b *bbrSender) checkIfFullBandwidthReached(lastPacketSendState *sendTimeState) { + if b.lastSampleIsAppLimited { + return + } + + target := Bandwidth(float64(b.bandwidthAtLastRound) * startupGrowthTarget) + if b.bandwidthEstimate() >= target { + b.bandwidthAtLastRound = b.bandwidthEstimate() + b.roundsWithoutBandwidthGain = 0 + if b.expireAckAggregationInStartup { + // Expire old excess delivery measurements now that bandwidth increased. + b.sampler.ResetMaxAckHeightTracker(0, b.roundTripCount) + } + return + } + + b.roundsWithoutBandwidthGain++ + if b.roundsWithoutBandwidthGain >= b.numStartupRtts || + b.shouldExitStartupDueToLoss(lastPacketSendState) { + b.isAtFullBandwidth = true + } +} + +func (b *bbrSender) maybeAppLimited(bytesInFlight congestion.ByteCount) { + if bytesInFlight < b.getTargetCongestionWindow(1) { + b.sampler.OnAppLimited() + } +} + +// Transitions from STARTUP to DRAIN and from DRAIN to PROBE_BW if +// appropriate. +func (b *bbrSender) maybeExitStartupOrDrain(now monotime.Time) { + if b.mode == bbrModeStartup && b.isAtFullBandwidth { + b.mode = bbrModeDrain + // b.maybeTraceStateChange(logging.CongestionStateDrain) + b.pacingGain = b.drainGain + b.congestionWindowGain = b.highCwndGain + + if b.debug { + b.debugPrint("Phase: DRAIN") + } + } + if b.mode == bbrModeDrain && b.bytesInFlight <= b.getTargetCongestionWindow(1) { + b.enterProbeBandwidthMode(now) + } +} + +// Decides whether to enter or exit PROBE_RTT. +func (b *bbrSender) maybeEnterOrExitProbeRtt(now monotime.Time, isRoundStart, minRttExpired bool) { + if minRttExpired && !b.exitingQuiescence && b.mode != bbrModeProbeRtt { + b.mode = bbrModeProbeRtt + // b.maybeTraceStateChange(logging.CongestionStateProbRtt) + b.pacingGain = 1.0 + // Do not decide on the time to exit PROBE_RTT until the |bytes_in_flight| + // is at the target small value. + b.exitProbeRttAt = 0 + + if b.debug { + b.debugPrint("BandwidthEstimate: %s, CongestionWindowGain: %.2f, PacingGain: %.2f, PacingRate: %s", + formatSpeed(b.bandwidthEstimate()), b.congestionWindowGain, b.pacingGain, formatSpeed(b.PacingRate())) + b.debugPrint("Phase: PROBE_RTT") + } + } + + if b.mode == bbrModeProbeRtt { + b.sampler.OnAppLimited() + // b.maybeTraceStateChange(logging.CongestionStateApplicationLimited) + + if b.exitProbeRttAt.IsZero() { + // If the window has reached the appropriate size, schedule exiting + // PROBE_RTT. The CWND during PROBE_RTT is kMinimumCongestionWindow, but + // we allow an extra packet since QUIC checks CWND before sending a + // packet. + if b.bytesInFlight < b.probeRttCongestionWindow()+congestion.MaxPacketBufferSize { + b.exitProbeRttAt = now.Add(probeRttTime) + b.probeRttRoundPassed = false + } + } else { + if isRoundStart { + b.probeRttRoundPassed = true + } + if now.Sub(b.exitProbeRttAt) >= 0 && b.probeRttRoundPassed { + b.minRttTimestamp = now + if b.debug { + b.debugPrint("MinRTT: %s", b.getMinRtt()) + } + if !b.isAtFullBandwidth { + b.enterStartupMode(now) + } else { + b.enterProbeBandwidthMode(now) + } + } + } + } + + b.exitingQuiescence = false +} + +// Determines whether BBR needs to enter, exit or advance state of the +// recovery. +func (b *bbrSender) updateRecoveryState(lastAckedPacket congestion.PacketNumber, hasLosses, isRoundStart bool) { + // Disable recovery in startup, if loss-based exit is enabled. + if !b.isAtFullBandwidth { + return + } + + // Exit recovery when there are no losses for a round. + if hasLosses { + b.endRecoveryAt = b.lastSentPacket + } + + switch b.recoveryState { + case bbrRecoveryStateNotInRecovery: + if hasLosses { + b.recoveryState = bbrRecoveryStateConservation + // This will cause the |recovery_window_| to be set to the correct + // value in CalculateRecoveryWindow(). + b.recoveryWindow = 0 + // Since the conservation phase is meant to be lasting for a whole + // round, extend the current round as if it were started right now. + b.currentRoundTripEnd = b.lastSentPacket + } + case bbrRecoveryStateConservation: + if isRoundStart { + b.recoveryState = bbrRecoveryStateGrowth + } + fallthrough + case bbrRecoveryStateGrowth: + // Exit recovery if appropriate. + if !hasLosses && lastAckedPacket > b.endRecoveryAt { + b.recoveryState = bbrRecoveryStateNotInRecovery + } + } +} + +// Determines the appropriate pacing rate for the connection. +func (b *bbrSender) calculatePacingRate(bytesLost congestion.ByteCount) { + if b.bandwidthEstimate() == 0 { + return + } + + targetRate := Bandwidth(b.pacingGain * float64(b.bandwidthEstimate())) + if b.isAtFullBandwidth { + b.pacingRate = targetRate + return + } + + // Pace at the rate of initial_window / RTT as soon as RTT measurements are + // available. + if b.pacingRate == 0 && b.rttStats.MinRTT() != 0 { + b.pacingRate = BandwidthFromDelta(b.initialCongestionWindow, b.rttStats.MinRTT()) + return + } + + if b.detectOvershooting { + b.bytesLostWhileDetectingOvershooting += bytesLost + // Check for overshooting with network parameters adjusted when pacing rate + // > target_rate and loss has been detected. + if b.pacingRate > targetRate && b.bytesLostWhileDetectingOvershooting > 0 { + if b.hasNoAppLimitedSample || + b.bytesLostWhileDetectingOvershooting*congestion.ByteCount(b.bytesLostMultiplierWhileDetectingOvershooting) > b.initialCongestionWindow { + // We are fairly sure overshoot happens if 1) there is at least one + // non app-limited bw sample or 2) half of IW gets lost. Slow pacing + // rate. + b.pacingRate = max(targetRate, BandwidthFromDelta(b.cwndToCalculateMinPacingRate, b.rttStats.MinRTT())) + b.bytesLostWhileDetectingOvershooting = 0 + b.detectOvershooting = false + } + } + } + + // Do not decrease the pacing rate during startup. + b.pacingRate = max(b.pacingRate, targetRate) +} + +// Determines the appropriate congestion window for the connection. +func (b *bbrSender) calculateCongestionWindow(bytesAcked, excessAcked congestion.ByteCount) { + if b.mode == bbrModeProbeRtt { + return + } + + targetWindow := b.getTargetCongestionWindow(b.congestionWindowGain) + if b.isAtFullBandwidth { + // Add the max recently measured ack aggregation to CWND. + targetWindow += b.sampler.MaxAckHeight() + } else if b.enableAckAggregationDuringStartup { + // Add the most recent excess acked. Because CWND never decreases in + // STARTUP, this will automatically create a very localized max filter. + targetWindow += excessAcked + } + + // Instead of immediately setting the target CWND as the new one, BBR grows + // the CWND towards |target_window| by only increasing it |bytes_acked| at a + // time. + if b.isAtFullBandwidth { + b.congestionWindow = min(targetWindow, b.congestionWindow+bytesAcked) + } else if b.congestionWindow < targetWindow || + b.sampler.TotalBytesAcked() < b.initialCongestionWindow { + // If the connection is not yet out of startup phase, do not decrease the + // window. + b.congestionWindow += bytesAcked + } + + // Enforce the limits on the congestion window. + b.congestionWindow = max(b.congestionWindow, b.minCongestionWindow) + b.congestionWindow = min(b.congestionWindow, b.maxCongestionWindow) +} + +// Determines the appropriate window that constrains the in-flight during recovery. +func (b *bbrSender) calculateRecoveryWindow(bytesAcked, bytesLost congestion.ByteCount) { + if b.recoveryState == bbrRecoveryStateNotInRecovery { + return + } + + // Set up the initial recovery window. + if b.recoveryWindow == 0 { + b.recoveryWindow = b.bytesInFlight + bytesAcked + b.recoveryWindow = max(b.minCongestionWindow, b.recoveryWindow) + return + } + + // Remove losses from the recovery window, while accounting for a potential + // integer underflow. + if b.recoveryWindow >= bytesLost { + b.recoveryWindow = b.recoveryWindow - bytesLost + } else { + b.recoveryWindow = b.maxDatagramSize + } + + // In CONSERVATION mode, just subtracting losses is sufficient. In GROWTH, + // release additional |bytes_acked| to achieve a slow-start-like behavior. + if b.recoveryState == bbrRecoveryStateGrowth { + b.recoveryWindow += bytesAcked + } + + // Always allow sending at least |bytes_acked| in response. + b.recoveryWindow = max(b.recoveryWindow, b.bytesInFlight+bytesAcked) + b.recoveryWindow = max(b.minCongestionWindow, b.recoveryWindow) +} + +// Return whether we should exit STARTUP due to excessive loss. +func (b *bbrSender) shouldExitStartupDueToLoss(lastPacketSendState *sendTimeState) bool { + if b.numLossEventsInRound < defaultStartupFullLossCount || !lastPacketSendState.isValid { + return false + } + + inflightAtSend := lastPacketSendState.bytesInFlight + + if inflightAtSend > 0 && b.bytesLostInRound > 0 { + if b.bytesLostInRound > congestion.ByteCount(float64(inflightAtSend)*quicBbr2DefaultLossThreshold) { + return true + } + return false + } + return false +} + +func (b *bbrSender) debugPrint(format string, a ...any) { + fmt.Printf("[BBRSender] [%s] %s\n", + time.Now().Format("15:04:05"), + fmt.Sprintf(format, a...)) +} + +func bdpFromRttAndBandwidth(rtt time.Duration, bandwidth Bandwidth) congestion.ByteCount { + return congestion.ByteCount(rtt) * congestion.ByteCount(bandwidth) / congestion.ByteCount(BytesPerSecond) / congestion.ByteCount(time.Second) +} + +func GetInitialPacketSize(addr net.Addr) congestion.ByteCount { + // If this is not a UDP address, we don't know anything about the MTU. + // Use the minimum size of an Initial packet as the max packet size. + if _, ok := addr.(*net.UDPAddr); ok { + return congestion.InitialPacketSize + } else { + return congestion.MinInitialPacketSize + } +} + +func formatSpeed(bw Bandwidth) string { + bwf := float64(bw) + units := []string{"bps", "Kbps", "Mbps", "Gbps"} + unitIndex := 0 + for bwf > 1000 && unitIndex < len(units)-1 { + bwf /= 1000 + unitIndex++ + } + return fmt.Sprintf("%.2f %s", bwf, units[unitIndex]) +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/bbr_sender_test.go b/third_party/hysteria-core/internal/congestion/bbr/bbr_sender_test.go new file mode 100644 index 0000000..e4f0259 --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/bbr_sender_test.go @@ -0,0 +1,226 @@ +package bbr + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/monotime" + "github.com/stretchr/testify/require" +) + +type fixedClock struct{ now monotime.Time } + +func (c fixedClock) Now() monotime.Time { return c.now } + +type fixedRTTStats struct { + min, latest, smoothed, deviation, maxAckDelay time.Duration +} + +func (s *fixedRTTStats) MinRTT() time.Duration { return s.min } +func (s *fixedRTTStats) LatestRTT() time.Duration { return s.latest } +func (s *fixedRTTStats) SmoothedRTT() time.Duration { return s.smoothed } +func (s *fixedRTTStats) MeanDeviation() time.Duration { return s.deviation } +func (s *fixedRTTStats) MaxAckDelay() time.Duration { return s.maxAckDelay } +func (s *fixedRTTStats) PTO(bool) time.Duration { return s.smoothed + 4*s.deviation } +func (s *fixedRTTStats) UpdateRTT(send, _ time.Duration) { s.latest, s.smoothed = send, send } +func (s *fixedRTTStats) SetMaxAckDelay(delay time.Duration) { s.maxAckDelay = delay } +func (s *fixedRTTStats) SetInitialRTT(rtt time.Duration) { s.min, s.latest, s.smoothed = rtt, rtt, rtt } + +func TestSetMaxDatagramSizeRescalesPacketSizedWindows(t *testing.T) { + const oldMaxDatagramSize = congestion.ByteCount(1000) + const newMaxDatagramSize = congestion.ByteCount(1400) + const initialCongestionWindowPackets = congestion.ByteCount(20) + const maxCongestionWindowPackets = congestion.ByteCount(80) + + b := newBbrSender( + DefaultClock{}, + oldMaxDatagramSize, + initialCongestionWindowPackets*oldMaxDatagramSize, + maxCongestionWindowPackets*oldMaxDatagramSize, + ProfileStandard, + ) + b.congestionWindow = b.initialCongestionWindow + + b.SetMaxDatagramSize(newMaxDatagramSize) + + require.Equal(t, initialCongestionWindowPackets*newMaxDatagramSize, b.initialCongestionWindow) + require.Equal(t, maxCongestionWindowPackets*newMaxDatagramSize, b.maxCongestionWindow) + require.Equal(t, minCongestionWindowPackets*newMaxDatagramSize, b.minCongestionWindow) + require.Equal(t, initialCongestionWindowPackets*newMaxDatagramSize, b.congestionWindow) +} + +func TestSetMaxDatagramSizeClampsCongestionWindow(t *testing.T) { + const oldMaxDatagramSize = congestion.ByteCount(1000) + const newMaxDatagramSize = congestion.ByteCount(1400) + + b := NewBbrSender(DefaultClock{}, oldMaxDatagramSize, ProfileStandard) + b.congestionWindow = b.minCongestionWindow + oldMaxDatagramSize + b.recoveryWindow = b.minCongestionWindow + oldMaxDatagramSize + + b.SetMaxDatagramSize(newMaxDatagramSize) + + require.Equal(t, b.minCongestionWindow, b.congestionWindow) + require.Equal(t, b.minCongestionWindow, b.recoveryWindow) +} + +func TestNewBbrSenderAppliesProfiles(t *testing.T) { + testCases := []struct { + name string + profile Profile + highGain float64 + highCwndGain float64 + congestionWindowGainConstant float64 + numStartupRtts int64 + drainToTarget bool + detectOvershooting bool + bytesLostMultiplier uint8 + enableAckAggregationDuringStartup bool + expireAckAggregationInStartup bool + enableOverestimateAvoidance bool + reduceExtraAckedOnBandwidthIncrease bool + }{ + { + name: "standard", + profile: ProfileStandard, + highGain: defaultHighGain, + highCwndGain: derivedHighCWNDGain, + congestionWindowGainConstant: 2.0, + numStartupRtts: roundTripsWithoutGrowthBeforeExitingStartup, + bytesLostMultiplier: 2, + }, + { + name: "conservative", + profile: ProfileConservative, + highGain: 2.25, + highCwndGain: 1.75, + congestionWindowGainConstant: 1.75, + numStartupRtts: 2, + drainToTarget: true, + detectOvershooting: true, + bytesLostMultiplier: 1, + enableOverestimateAvoidance: true, + reduceExtraAckedOnBandwidthIncrease: true, + }, + { + name: "aggressive", + profile: ProfileAggressive, + highGain: 3.0, + highCwndGain: 2.25, + congestionWindowGainConstant: 2.5, + numStartupRtts: 4, + bytesLostMultiplier: 2, + enableAckAggregationDuringStartup: true, + expireAckAggregationInStartup: true, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + b := NewBbrSender(DefaultClock{}, congestion.InitialPacketSize, tc.profile) + require.Equal(t, tc.profile, b.profile) + require.Equal(t, tc.highGain, b.highGain) + require.Equal(t, tc.highCwndGain, b.highCwndGain) + require.Equal(t, tc.congestionWindowGainConstant, b.congestionWindowGainConstant) + require.Equal(t, tc.numStartupRtts, b.numStartupRtts) + require.Equal(t, tc.drainToTarget, b.drainToTarget) + require.Equal(t, tc.detectOvershooting, b.detectOvershooting) + require.Equal(t, tc.bytesLostMultiplier, b.bytesLostMultiplierWhileDetectingOvershooting) + require.Equal(t, tc.enableAckAggregationDuringStartup, b.enableAckAggregationDuringStartup) + require.Equal(t, tc.expireAckAggregationInStartup, b.expireAckAggregationInStartup) + require.Equal(t, tc.enableOverestimateAvoidance, b.sampler.IsOverestimateAvoidanceEnabled()) + require.Equal(t, tc.reduceExtraAckedOnBandwidthIncrease, b.sampler.maxAckHeightTracker.reduceExtraAckedOnBandwidthIncrease) + require.Equal(t, b.highGain, b.pacingGain) + require.Equal(t, b.highCwndGain, b.congestionWindowGain) + }) + } +} + +func TestParseProfile(t *testing.T) { + profile, err := ParseProfile("") + require.NoError(t, err) + require.Equal(t, ProfileStandard, profile) + + profile, err = ParseProfile("Aggressive") + require.NoError(t, err) + require.Equal(t, ProfileAggressive, profile) + + _, err = ParseProfile("turbo") + require.EqualError(t, err, `unsupported BBR profile "turbo"`) +} + +func TestCongestionEventUpdatesDeliveryRateAndMinimumRTT(t *testing.T) { + now := monotime.Now() + b := NewBbrSender(fixedClock{now: now}, 1200, ProfileStandard) + b.SetRTTStatsProvider(&fixedRTTStats{min: 100 * time.Millisecond, smoothed: 100 * time.Millisecond}) + + const packetSize = congestion.ByteCount(1200) + b.OnPacketSent(now, 0, 1, packetSize, true) + ackedAt := now.Add(100 * time.Millisecond) + b.OnCongestionEventEx(packetSize, ackedAt, []congestion.AckedPacketInfo{{ + PacketNumber: 1, + BytesAcked: packetSize, + ReceivedTime: ackedAt, + }}, nil) + + require.Equal(t, 100*time.Millisecond, b.minRtt) + require.Positive(t, b.bandwidthEstimate()) + require.Equal(t, packetSize, b.sampler.TotalBytesAcked()) + require.Positive(t, b.PacingRate()) +} + +func TestBBRStateMachineStartupDrainProbeBandwidthAndProbeRTT(t *testing.T) { + now := monotime.Now() + b := NewBbrSender(fixedClock{now: now}, 1200, ProfileStandard) + b.SetRTTStatsProvider(&fixedRTTStats{min: 100 * time.Millisecond, smoothed: 100 * time.Millisecond}) + b.minRtt = 100 * time.Millisecond + b.minRttTimestamp = now + b.maxBandwidth.Update(BandwidthFromDelta(12000, 100*time.Millisecond), 1) + b.isAtFullBandwidth = true + + target := b.getTargetCongestionWindow(1) + b.bytesInFlight = target + b.maxDatagramSize + b.maybeExitStartupOrDrain(now) + require.EqualValues(t, bbrModeDrain, b.mode) + require.Equal(t, b.drainGain, b.pacingGain) + + b.bytesInFlight = target + b.maybeExitStartupOrDrain(now.Add(time.Millisecond)) + require.EqualValues(t, bbrModeProbeBw, b.mode) + require.Equal(t, b.congestionWindowGainConstant, b.congestionWindowGain) + + b.exitingQuiescence = false + b.bytesInFlight = 0 + probeStart := now.Add(minRttExpiry + time.Second) + b.maybeEnterOrExitProbeRtt(probeStart, true, true) + require.EqualValues(t, bbrModeProbeRtt, b.mode) + require.Equal(t, b.minCongestionWindow, b.GetCongestionWindow()) + require.Equal(t, probeStart.Add(probeRttTime), b.exitProbeRttAt) + + probeEnd := probeStart.Add(probeRttTime + time.Millisecond) + b.maybeEnterOrExitProbeRtt(probeEnd, true, false) + require.EqualValues(t, bbrModeProbeBw, b.mode) + require.Equal(t, probeEnd, b.minRttTimestamp) +} + +func TestBBRLossRecoveryConservationGrowthAndExit(t *testing.T) { + b := NewBbrSender(DefaultClock{}, 1200, ProfileStandard) + b.isAtFullBandwidth = true + b.lastSentPacket = 20 + b.bytesInFlight = 12000 + + b.updateRecoveryState(10, true, false) + require.EqualValues(t, bbrRecoveryStateConservation, b.recoveryState) + require.Equal(t, congestion.PacketNumber(20), b.endRecoveryAt) + b.calculateRecoveryWindow(1200, 1200) + require.GreaterOrEqual(t, b.recoveryWindow, b.minCongestionWindow) + + b.updateRecoveryState(20, false, true) + require.EqualValues(t, bbrRecoveryStateGrowth, b.recoveryState) + before := b.recoveryWindow + b.calculateRecoveryWindow(1200, 0) + require.Greater(t, b.recoveryWindow, before) + + b.updateRecoveryState(21, false, false) + require.EqualValues(t, bbrRecoveryStateNotInRecovery, b.recoveryState) +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/clock.go b/third_party/hysteria-core/internal/congestion/bbr/clock.go new file mode 100644 index 0000000..541987e --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/clock.go @@ -0,0 +1,18 @@ +package bbr + +import "github.com/apernet/quic-go/monotime" + +// A Clock returns the current time +type Clock interface { + Now() monotime.Time +} + +// DefaultClock implements the Clock interface using the Go stdlib clock. +type DefaultClock struct{} + +var _ Clock = DefaultClock{} + +// Now gets the current time +func (DefaultClock) Now() monotime.Time { + return monotime.Now() +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/packet_number_indexed_queue.go b/third_party/hysteria-core/internal/congestion/bbr/packet_number_indexed_queue.go new file mode 100644 index 0000000..08b99de --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/packet_number_indexed_queue.go @@ -0,0 +1,199 @@ +package bbr + +import ( + "github.com/apernet/quic-go/congestion" +) + +// packetNumberIndexedQueue is a queue of mostly continuous numbered entries +// which supports the following operations: +// - adding elements to the end of the queue, or at some point past the end +// - removing elements in any order +// - retrieving elements +// If all elements are inserted in order, all of the operations above are +// amortized O(1) time. +// +// Internally, the data structure is a deque where each element is marked as +// present or not. The deque starts at the lowest present index. Whenever an +// element is removed, it's marked as not present, and the front of the deque is +// cleared of elements that are not present. +// +// The tail of the queue is not cleared due to the assumption of entries being +// inserted in order, though removing all elements of the queue will return it +// to its initial state. +// +// Note that this data structure is inherently hazardous, since an addition of +// just two entries will cause it to consume all of the memory available. +// Because of that, it is not a general-purpose container and should not be used +// as one. + +type entryWrapper[T any] struct { + present bool + entry T +} + +type packetNumberIndexedQueue[T any] struct { + entries RingBuffer[entryWrapper[T]] + numberOfPresentEntries int + firstPacket congestion.PacketNumber +} + +func newPacketNumberIndexedQueue[T any](size int) *packetNumberIndexedQueue[T] { + q := &packetNumberIndexedQueue[T]{ + firstPacket: invalidPacketNumber, + } + + q.entries.Init(size) + + return q +} + +// Emplace inserts data associated |packet_number| into (or past) the end of the +// queue, filling up the missing intermediate entries as necessary. Returns +// true if the element has been inserted successfully, false if it was already +// in the queue or inserted out of order. +func (p *packetNumberIndexedQueue[T]) Emplace(packetNumber congestion.PacketNumber, entry *T) bool { + if packetNumber == invalidPacketNumber || entry == nil { + return false + } + + if p.IsEmpty() { + p.entries.PushBack(entryWrapper[T]{ + present: true, + entry: *entry, + }) + p.numberOfPresentEntries = 1 + p.firstPacket = packetNumber + return true + } + + // Do not allow insertion out-of-order. + if packetNumber <= p.LastPacket() { + return false + } + + // Handle potentially missing elements. + offset := int(packetNumber - p.FirstPacket()) + if gap := offset - p.entries.Len(); gap > 0 { + for i := 0; i < gap; i++ { + p.entries.PushBack(entryWrapper[T]{}) + } + } + + p.entries.PushBack(entryWrapper[T]{ + present: true, + entry: *entry, + }) + p.numberOfPresentEntries++ + return true +} + +// GetEntry Retrieve the entry associated with the packet number. Returns the pointer +// to the entry in case of success, or nullptr if the entry does not exist. +func (p *packetNumberIndexedQueue[T]) GetEntry(packetNumber congestion.PacketNumber) *T { + ew := p.getEntryWraper(packetNumber) + if ew == nil { + return nil + } + + return &ew.entry +} + +// Remove, Same as above, but if an entry is present in the queue, also call f(entry) +// before removing it. +func (p *packetNumberIndexedQueue[T]) Remove(packetNumber congestion.PacketNumber, f func(T)) bool { + ew := p.getEntryWraper(packetNumber) + if ew == nil { + return false + } + if f != nil { + f(ew.entry) + } + ew.present = false + p.numberOfPresentEntries-- + + if packetNumber == p.FirstPacket() { + p.clearup() + } + + return true +} + +// RemoveUpTo, but not including |packet_number|. +// Unused slots in the front are also removed, which means when the function +// returns, |first_packet()| can be larger than |packet_number|. +func (p *packetNumberIndexedQueue[T]) RemoveUpTo(packetNumber congestion.PacketNumber) { + for !p.entries.Empty() && + p.firstPacket != invalidPacketNumber && + p.firstPacket < packetNumber { + if p.entries.Front().present { + p.numberOfPresentEntries-- + } + p.entries.PopFront() + p.firstPacket++ + } + p.clearup() + + return +} + +// IsEmpty return if queue is empty. +func (p *packetNumberIndexedQueue[T]) IsEmpty() bool { + return p.numberOfPresentEntries == 0 +} + +// NumberOfPresentEntries returns the number of entries in the queue. +func (p *packetNumberIndexedQueue[T]) NumberOfPresentEntries() int { + return p.numberOfPresentEntries +} + +// EntrySlotsUsed returns the number of entries allocated in the underlying deque. This is +// proportional to the memory usage of the queue. +func (p *packetNumberIndexedQueue[T]) EntrySlotsUsed() int { + return p.entries.Len() +} + +// FirstPacket returns packet number of the first entry in the queue. +func (p *packetNumberIndexedQueue[T]) FirstPacket() (packetNumber congestion.PacketNumber) { + return p.firstPacket +} + +// LastPacket returns packet number of the last entry ever inserted in the queue. Note that the +// entry in question may have already been removed. Zero if the queue is +// empty. +func (p *packetNumberIndexedQueue[T]) LastPacket() (packetNumber congestion.PacketNumber) { + if p.IsEmpty() { + return invalidPacketNumber + } + + return p.firstPacket + congestion.PacketNumber(p.entries.Len()-1) +} + +func (p *packetNumberIndexedQueue[T]) clearup() { + for !p.entries.Empty() && !p.entries.Front().present { + p.entries.PopFront() + p.firstPacket++ + } + if p.entries.Empty() { + p.firstPacket = invalidPacketNumber + } +} + +func (p *packetNumberIndexedQueue[T]) getEntryWraper(packetNumber congestion.PacketNumber) *entryWrapper[T] { + if packetNumber == invalidPacketNumber || + p.IsEmpty() || + packetNumber < p.firstPacket { + return nil + } + + offset := int(packetNumber - p.firstPacket) + if offset >= p.entries.Len() { + return nil + } + + ew := p.entries.Offset(offset) + if ew == nil || !ew.present { + return nil + } + + return ew +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/ringbuffer.go b/third_party/hysteria-core/internal/congestion/bbr/ringbuffer.go new file mode 100644 index 0000000..ed92d4c --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/ringbuffer.go @@ -0,0 +1,118 @@ +package bbr + +// A RingBuffer is a ring buffer. +// It acts as a heap that doesn't cause any allocations. +type RingBuffer[T any] struct { + ring []T + headPos, tailPos int + full bool +} + +// Init preallocs a buffer with a certain size. +func (r *RingBuffer[T]) Init(size int) { + r.ring = make([]T, size) +} + +// Len returns the number of elements in the ring buffer. +func (r *RingBuffer[T]) Len() int { + if r.full { + return len(r.ring) + } + if r.tailPos >= r.headPos { + return r.tailPos - r.headPos + } + return r.tailPos - r.headPos + len(r.ring) +} + +// Empty says if the ring buffer is empty. +func (r *RingBuffer[T]) Empty() bool { + return !r.full && r.headPos == r.tailPos +} + +// PushBack adds a new element. +// If the ring buffer is full, its capacity is increased first. +func (r *RingBuffer[T]) PushBack(t T) { + if r.full || len(r.ring) == 0 { + r.grow() + } + r.ring[r.tailPos] = t + r.tailPos++ + if r.tailPos == len(r.ring) { + r.tailPos = 0 + } + if r.tailPos == r.headPos { + r.full = true + } +} + +// PopFront returns the next element. +// It must not be called when the buffer is empty, that means that +// callers might need to check if there are elements in the buffer first. +func (r *RingBuffer[T]) PopFront() T { + if r.Empty() { + panic("github.com/quic-go/quic-go/internal/utils/ringbuffer: pop from an empty queue") + } + r.full = false + t := r.ring[r.headPos] + r.ring[r.headPos] = *new(T) + r.headPos++ + if r.headPos == len(r.ring) { + r.headPos = 0 + } + return t +} + +// Offset returns the offset element. +// It must not be called when the buffer is empty, that means that +// callers might need to check if there are elements in the buffer first +// and check if the index larger than buffer length. +func (r *RingBuffer[T]) Offset(index int) *T { + if r.Empty() || index >= r.Len() { + panic("github.com/quic-go/quic-go/internal/utils/ringbuffer: offset from invalid index") + } + offset := (r.headPos + index) % len(r.ring) + return &r.ring[offset] +} + +// Front returns the front element. +// It must not be called when the buffer is empty, that means that +// callers might need to check if there are elements in the buffer first. +func (r *RingBuffer[T]) Front() *T { + if r.Empty() { + panic("github.com/quic-go/quic-go/internal/utils/ringbuffer: front from an empty queue") + } + return &r.ring[r.headPos] +} + +// Back returns the back element. +// It must not be called when the buffer is empty, that means that +// callers might need to check if there are elements in the buffer first. +func (r *RingBuffer[T]) Back() *T { + if r.Empty() { + panic("github.com/quic-go/quic-go/internal/utils/ringbuffer: back from an empty queue") + } + return r.Offset(r.Len() - 1) +} + +// Grow the maximum size of the queue. +// This method assume the queue is full. +func (r *RingBuffer[T]) grow() { + oldRing := r.ring + newSize := len(oldRing) * 2 + if newSize == 0 { + newSize = 1 + } + r.ring = make([]T, newSize) + headLen := copy(r.ring, oldRing[r.headPos:]) + copy(r.ring[headLen:], oldRing[:r.headPos]) + r.headPos, r.tailPos, r.full = 0, len(oldRing), false +} + +// Clear removes all elements. +func (r *RingBuffer[T]) Clear() { + var zeroValue T + for i := range r.ring { + r.ring[i] = zeroValue + } + r.headPos, r.tailPos, r.full = 0, 0, false +} diff --git a/third_party/hysteria-core/internal/congestion/bbr/windowed_filter.go b/third_party/hysteria-core/internal/congestion/bbr/windowed_filter.go new file mode 100644 index 0000000..4773bce --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/bbr/windowed_filter.go @@ -0,0 +1,162 @@ +package bbr + +import ( + "golang.org/x/exp/constraints" +) + +// Implements Kathleen Nichols' algorithm for tracking the minimum (or maximum) +// estimate of a stream of samples over some fixed time interval. (E.g., +// the minimum RTT over the past five minutes.) The algorithm keeps track of +// the best, second best, and third best min (or max) estimates, maintaining an +// invariant that the measurement time of the n'th best >= n-1'th best. + +// The algorithm works as follows. On a reset, all three estimates are set to +// the same sample. The second best estimate is then recorded in the second +// quarter of the window, and a third best estimate is recorded in the second +// half of the window, bounding the worst case error when the true min is +// monotonically increasing (or true max is monotonically decreasing) over the +// window. +// +// A new best sample replaces all three estimates, since the new best is lower +// (or higher) than everything else in the window and it is the most recent. +// The window thus effectively gets reset on every new min. The same property +// holds true for second best and third best estimates. Specifically, when a +// sample arrives that is better than the second best but not better than the +// best, it replaces the second and third best estimates but not the best +// estimate. Similarly, a sample that is better than the third best estimate +// but not the other estimates replaces only the third best estimate. +// +// Finally, when the best expires, it is replaced by the second best, which in +// turn is replaced by the third best. The newest sample replaces the third +// best. + +type WindowedFilterValue interface { + any +} + +type WindowedFilterTime interface { + constraints.Integer | constraints.Float +} + +type WindowedFilter[V WindowedFilterValue, T WindowedFilterTime] struct { + // Time length of window. + windowLength T + estimates []entry[V, T] + comparator func(V, V) int +} + +type entry[V WindowedFilterValue, T WindowedFilterTime] struct { + sample V + time T +} + +// Compares two values and returns true if the first is greater than or equal +// to the second. +func MaxFilter[O constraints.Ordered](a, b O) int { + if a > b { + return 1 + } else if a < b { + return -1 + } + return 0 +} + +// Compares two values and returns true if the first is less than or equal +// to the second. +func MinFilter[O constraints.Ordered](a, b O) int { + if a < b { + return 1 + } else if a > b { + return -1 + } + return 0 +} + +func NewWindowedFilter[V WindowedFilterValue, T WindowedFilterTime](windowLength T, comparator func(V, V) int) *WindowedFilter[V, T] { + return &WindowedFilter[V, T]{ + windowLength: windowLength, + estimates: make([]entry[V, T], 3, 3), + comparator: comparator, + } +} + +// Changes the window length. Does not update any current samples. +func (f *WindowedFilter[V, T]) SetWindowLength(windowLength T) { + f.windowLength = windowLength +} + +func (f *WindowedFilter[V, T]) GetBest() V { + return f.estimates[0].sample +} + +func (f *WindowedFilter[V, T]) GetSecondBest() V { + return f.estimates[1].sample +} + +func (f *WindowedFilter[V, T]) GetThirdBest() V { + return f.estimates[2].sample +} + +// Updates best estimates with |sample|, and expires and updates best +// estimates as necessary. +func (f *WindowedFilter[V, T]) Update(newSample V, newTime T) { + // Reset all estimates if they have not yet been initialized, if new sample + // is a new best, or if the newest recorded estimate is too old. + if f.comparator(f.estimates[0].sample, *new(V)) == 0 || + f.comparator(newSample, f.estimates[0].sample) >= 0 || + newTime-f.estimates[2].time > f.windowLength { + f.Reset(newSample, newTime) + return + } + + if f.comparator(newSample, f.estimates[1].sample) >= 0 { + f.estimates[1] = entry[V, T]{newSample, newTime} + f.estimates[2] = f.estimates[1] + } else if f.comparator(newSample, f.estimates[2].sample) >= 0 { + f.estimates[2] = entry[V, T]{newSample, newTime} + } + + // Expire and update estimates as necessary. + if newTime-f.estimates[0].time > f.windowLength { + // The best estimate hasn't been updated for an entire window, so promote + // second and third best estimates. + f.estimates[0] = f.estimates[1] + f.estimates[1] = f.estimates[2] + f.estimates[2] = entry[V, T]{newSample, newTime} + // Need to iterate one more time. Check if the new best estimate is + // outside the window as well, since it may also have been recorded a + // long time ago. Don't need to iterate once more since we cover that + // case at the beginning of the method. + if newTime-f.estimates[0].time > f.windowLength { + f.estimates[0] = f.estimates[1] + f.estimates[1] = f.estimates[2] + } + return + } + if f.comparator(f.estimates[1].sample, f.estimates[0].sample) == 0 && + newTime-f.estimates[1].time > f.windowLength/4 { + // A quarter of the window has passed without a better sample, so the + // second-best estimate is taken from the second quarter of the window. + f.estimates[1] = entry[V, T]{newSample, newTime} + f.estimates[2] = f.estimates[1] + return + } + + if f.comparator(f.estimates[2].sample, f.estimates[1].sample) == 0 && + newTime-f.estimates[2].time > f.windowLength/2 { + // We've passed a half of the window without a better estimate, so take + // a third-best estimate from the second half of the window. + f.estimates[2] = entry[V, T]{newSample, newTime} + } +} + +// Resets all estimates to new sample. +func (f *WindowedFilter[V, T]) Reset(newSample V, newTime T) { + f.estimates[2] = entry[V, T]{newSample, newTime} + f.estimates[1] = f.estimates[2] + f.estimates[0] = f.estimates[1] +} + +func (f *WindowedFilter[V, T]) Clear() { + f.estimates = make([]entry[V, T], 3, 3) +} diff --git a/third_party/hysteria-core/internal/congestion/brutal/brutal.go b/third_party/hysteria-core/internal/congestion/brutal/brutal.go new file mode 100644 index 0000000..ec61bf1 --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/brutal/brutal.go @@ -0,0 +1,193 @@ +package brutal + +import ( + "fmt" + "os" + "strconv" + "time" + + "github.com/apernet/hysteria/core/v2/internal/congestion/common" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/monotime" +) + +const ( + pktInfoSlotCount = 5 // slot index is based on seconds, so this is basically how many seconds we sample + minSampleCount = 50 + minAckRate = 0.8 + congestionWindowMultiplier = 2 + + debugEnv = "HYSTERIA_BRUTAL_DEBUG" + debugPrintInterval = 2 +) + +var _ congestion.CongestionControl = &BrutalSender{} + +type BrutalSender struct { + rttStats congestion.RTTStatsProvider + bps congestion.ByteCount + maxDatagramSize congestion.ByteCount + pacer *common.Pacer + + pktInfoSlots [pktInfoSlotCount]pktInfo + ackRate float64 + + disableLossCompensation bool + + debug bool + lastAckPrintTimestamp int64 +} + +type pktInfo struct { + Timestamp int64 + AckCount uint64 + LossCount uint64 +} + +func NewBrutalSender(bps uint64, disableLossCompensation bool) *BrutalSender { + debug, _ := strconv.ParseBool(os.Getenv(debugEnv)) + bs := &BrutalSender{ + bps: congestion.ByteCount(bps), + maxDatagramSize: congestion.InitialPacketSize, + ackRate: 1, + disableLossCompensation: disableLossCompensation, + debug: debug, + } + bs.pacer = common.NewPacer(func() congestion.ByteCount { + return congestion.ByteCount(float64(bs.bps) / bs.ackRate) + }) + return bs +} + +func (b *BrutalSender) SetRTTStatsProvider(rttStats congestion.RTTStatsProvider) { + b.rttStats = rttStats +} + +func (b *BrutalSender) TimeUntilSend(bytesInFlight congestion.ByteCount) monotime.Time { + return b.pacer.TimeUntilSend() +} + +func (b *BrutalSender) HasPacingBudget(now monotime.Time) bool { + return b.pacer.Budget(now) >= b.maxDatagramSize +} + +func (b *BrutalSender) CanSend(bytesInFlight congestion.ByteCount) bool { + return bytesInFlight <= b.GetCongestionWindow() +} + +func (b *BrutalSender) GetCongestionWindow() congestion.ByteCount { + rtt := b.rttStats.SmoothedRTT() + if rtt <= 0 { + return 10240 + } + cwnd := congestion.ByteCount(float64(b.bps) * rtt.Seconds() * congestionWindowMultiplier / b.ackRate) + if cwnd < b.maxDatagramSize { + cwnd = b.maxDatagramSize + } + return cwnd +} + +func (b *BrutalSender) OnPacketSent(sentTime monotime.Time, bytesInFlight congestion.ByteCount, + packetNumber congestion.PacketNumber, bytes congestion.ByteCount, isRetransmittable bool, +) { + b.pacer.SentPacket(sentTime, bytes) +} + +func (b *BrutalSender) OnPacketAcked(number congestion.PacketNumber, ackedBytes congestion.ByteCount, + priorInFlight congestion.ByteCount, eventTime monotime.Time, +) { + // Stub +} + +func (b *BrutalSender) OnCongestionEvent(number congestion.PacketNumber, lostBytes congestion.ByteCount, + priorInFlight congestion.ByteCount, +) { + // Stub +} + +func (b *BrutalSender) OnCongestionEventEx(priorInFlight congestion.ByteCount, eventTime monotime.Time, ackedPackets []congestion.AckedPacketInfo, lostPackets []congestion.LostPacketInfo) { + currentTimestamp := int64(time.Duration(eventTime) / time.Second) + slot := currentTimestamp % pktInfoSlotCount + if b.pktInfoSlots[slot].Timestamp == currentTimestamp { + b.pktInfoSlots[slot].LossCount += uint64(len(lostPackets)) + b.pktInfoSlots[slot].AckCount += uint64(len(ackedPackets)) + } else { + // uninitialized slot or too old, reset + b.pktInfoSlots[slot].Timestamp = currentTimestamp + b.pktInfoSlots[slot].AckCount = uint64(len(ackedPackets)) + b.pktInfoSlots[slot].LossCount = uint64(len(lostPackets)) + } + b.updateAckRate(currentTimestamp) +} + +func (b *BrutalSender) SetMaxDatagramSize(size congestion.ByteCount) { + b.maxDatagramSize = size + b.pacer.SetMaxDatagramSize(size) + if b.debug { + b.debugPrint("SetMaxDatagramSize: %d", size) + } +} + +func (b *BrutalSender) updateAckRate(currentTimestamp int64) { + if b.disableLossCompensation { + b.ackRate = 1 + return + } + minTimestamp := currentTimestamp - pktInfoSlotCount + var ackCount, lossCount uint64 + for _, info := range b.pktInfoSlots { + if info.Timestamp < minTimestamp { + continue + } + ackCount += info.AckCount + lossCount += info.LossCount + } + if ackCount+lossCount < minSampleCount { + b.ackRate = 1 + if b.canPrintAckRate(currentTimestamp) { + b.lastAckPrintTimestamp = currentTimestamp + b.debugPrint("Not enough samples (total=%d, ack=%d, loss=%d, rtt=%d)", + ackCount+lossCount, ackCount, lossCount, b.rttStats.SmoothedRTT().Milliseconds()) + } + return + } + rate := float64(ackCount) / float64(ackCount+lossCount) + if rate < minAckRate { + b.ackRate = minAckRate + if b.canPrintAckRate(currentTimestamp) { + b.lastAckPrintTimestamp = currentTimestamp + b.debugPrint("ACK rate too low: %.2f, clamped to %.2f (total=%d, ack=%d, loss=%d, rtt=%d)", + rate, minAckRate, ackCount+lossCount, ackCount, lossCount, b.rttStats.SmoothedRTT().Milliseconds()) + } + return + } + b.ackRate = rate + if b.canPrintAckRate(currentTimestamp) { + b.lastAckPrintTimestamp = currentTimestamp + b.debugPrint("ACK rate: %.2f (total=%d, ack=%d, loss=%d, rtt=%d)", + rate, ackCount+lossCount, ackCount, lossCount, b.rttStats.SmoothedRTT().Milliseconds()) + } +} + +func (b *BrutalSender) InSlowStart() bool { + return false +} + +func (b *BrutalSender) InRecovery() bool { + return false +} + +func (b *BrutalSender) MaybeExitSlowStart() {} + +func (b *BrutalSender) OnRetransmissionTimeout(packetsRetransmitted bool) {} + +func (b *BrutalSender) canPrintAckRate(currentTimestamp int64) bool { + return b.debug && currentTimestamp-b.lastAckPrintTimestamp >= debugPrintInterval +} + +func (b *BrutalSender) debugPrint(format string, a ...any) { + fmt.Printf("[BrutalSender] [%s] %s\n", + time.Now().Format("15:04:05"), + fmt.Sprintf(format, a...)) +} diff --git a/third_party/hysteria-core/internal/congestion/brutal/brutal_test.go b/third_party/hysteria-core/internal/congestion/brutal/brutal_test.go new file mode 100644 index 0000000..2565beb --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/brutal/brutal_test.go @@ -0,0 +1,45 @@ +package brutal + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/monotime" +) + +// feedAckRate drives a single sampling slot with the given number of acked and +// lost packets and returns the resulting ackRate. +func feedAckRate(disableLossCompensation bool, ackCount, lossCount int) float64 { + b := NewBrutalSender(1000000, disableLossCompensation) + acked := make([]congestion.AckedPacketInfo, ackCount) + lost := make([]congestion.LostPacketInfo, lossCount) + // eventTime lands in a fixed slot; a single event carries enough samples. + b.OnCongestionEventEx(0, monotime.Time(5*time.Second), acked, lost) + return b.ackRate +} + +func TestBrutalLossCompensation(t *testing.T) { + tests := []struct { + name string + ack, loss int + want float64 // expected ackRate when compensation is ENABLED + }{ + {"no loss", 100, 0, 1.0}, + {"20% loss", 80, 20, 0.8}, + {"50% loss clamps to floor", 50, 50, minAckRate}, // 0.5 clamped up to 0.8 + {"few samples stays 1", 10, 5, 1.0}, // below minSampleCount + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Compensation enabled (default behavior): ackRate reacts to loss. + if got := feedAckRate(false, tt.ack, tt.loss); got != tt.want { + t.Errorf("compensation on: ackRate = %v, want %v", got, tt.want) + } + // Compensation disabled: ackRate must stay pinned at 1 regardless. + if got := feedAckRate(true, tt.ack, tt.loss); got != 1.0 { + t.Errorf("compensation off: ackRate = %v, want 1.0", got) + } + }) + } +} diff --git a/third_party/hysteria-core/internal/congestion/common/pacer.go b/third_party/hysteria-core/internal/congestion/common/pacer.go new file mode 100644 index 0000000..e0bddbe --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/common/pacer.go @@ -0,0 +1,80 @@ +package common + +import ( + "time" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/monotime" +) + +const ( + maxBurstPackets = 10 + maxBurstPacingDelayMultiplier = 4 +) + +// Pacer implements a token bucket pacing algorithm. +type Pacer struct { + budgetAtLastSent congestion.ByteCount + maxDatagramSize congestion.ByteCount + lastSentTime monotime.Time + getBandwidth func() congestion.ByteCount // in bytes/s +} + +func NewPacer(getBandwidth func() congestion.ByteCount) *Pacer { + p := &Pacer{ + budgetAtLastSent: maxBurstPackets * congestion.InitialPacketSize, + maxDatagramSize: congestion.InitialPacketSize, + getBandwidth: getBandwidth, + } + return p +} + +func (p *Pacer) SentPacket(sendTime monotime.Time, size congestion.ByteCount) { + budget := p.Budget(sendTime) + if size > budget { + p.budgetAtLastSent = 0 + } else { + p.budgetAtLastSent = budget - size + } + p.lastSentTime = sendTime +} + +func (p *Pacer) Budget(now monotime.Time) congestion.ByteCount { + if p.lastSentTime.IsZero() { + return p.maxBurstSize() + } + budget := p.budgetAtLastSent + (p.getBandwidth()*congestion.ByteCount(now.Sub(p.lastSentTime).Nanoseconds()))/1e9 + if budget < 0 { // protect against overflows + budget = congestion.ByteCount(1<<62 - 1) + } + return min(p.maxBurstSize(), budget) +} + +func (p *Pacer) maxBurstSize() congestion.ByteCount { + return max( + congestion.ByteCount((maxBurstPacingDelayMultiplier*congestion.MinPacingDelay).Nanoseconds())*p.getBandwidth()/1e9, + maxBurstPackets*p.maxDatagramSize, + ) +} + +// TimeUntilSend returns when the next packet should be sent. +// It returns the zero value if a packet can be sent immediately. +func (p *Pacer) TimeUntilSend() monotime.Time { + if p.budgetAtLastSent >= p.maxDatagramSize { + return 0 + } + diff := 1e9 * uint64(p.maxDatagramSize-p.budgetAtLastSent) + bw := uint64(p.getBandwidth()) + // We might need to round up this value. + // Otherwise, we might have a budget (slightly) smaller than the datagram size when the timer expires. + d := diff / bw + // this is effectively a math.Ceil, but using only integer math + if diff%bw > 0 { + d++ + } + return p.lastSentTime.Add(max(congestion.MinPacingDelay, time.Duration(d)*time.Nanosecond)) +} + +func (p *Pacer) SetMaxDatagramSize(s congestion.ByteCount) { + p.maxDatagramSize = s +} diff --git a/third_party/hysteria-core/internal/congestion/utils.go b/third_party/hysteria-core/internal/congestion/utils.go new file mode 100644 index 0000000..1d72811 --- /dev/null +++ b/third_party/hysteria-core/internal/congestion/utils.go @@ -0,0 +1,72 @@ +package congestion + +import ( + "fmt" + "strings" + + "github.com/apernet/hysteria/core/v2/internal/congestion/bbr" + "github.com/apernet/hysteria/core/v2/internal/congestion/brutal" + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/congestion" +) + +const ( + TypeBBR = "bbr" + TypeReno = "reno" +) + +func NormalizeType(congestionType string) (string, error) { + switch normalized := strings.ToLower(congestionType); normalized { + case "", TypeBBR: + return TypeBBR, nil + case TypeReno: + return TypeReno, nil + default: + return "", fmt.Errorf("unsupported congestion type %q", congestionType) + } +} + +func NormalizeBBRProfile(profile string) (string, error) { + normalized, err := bbr.ParseProfile(profile) + if err != nil { + return "", err + } + return string(normalized), nil +} + +func UseBBR(conn *quic.Conn, profile bbr.Profile) { + conn.SetCongestionControl(bbr.NewBbrSender( + bbr.DefaultClock{}, + seedPacketSize(conn.InitialPacketSize(), bbr.GetInitialPacketSize(conn.RemoteAddr())), + profile, + )) +} + +// seedPacketSize picks the datagram size to seed a replacement congestion +// controller with, given the size QUIC itself starts at and the guess derived +// from the remote address. +// +// The seed must not exceed what QUIC actually starts at. If it does, the first +// path MTU probe can land between the two: QUIC sees an increase and reports +// it, but the controller sees a decrease, which it cannot represent. Taking the +// smaller of the two keeps the address-based guess as a floor for connections +// whose path we can't reason about, while never seeding above QUIC. +func seedPacketSize(quicSize, byAddr congestion.ByteCount) congestion.ByteCount { + if quicSize <= 0 { + return byAddr + } + return min(quicSize, byAddr) +} + +func UseBrutal(conn *quic.Conn, tx uint64, disableLossCompensation bool) { + conn.SetCongestionControl(brutal.NewBrutalSender(tx, disableLossCompensation)) +} + +func UseConfigured(conn *quic.Conn, congestionType, bbrProfile string) { + switch congestionType { + case TypeReno: + return + default: + UseBBR(conn, bbr.Profile(bbrProfile)) + } +} diff --git a/third_party/hysteria-core/internal/frag/frag.go b/third_party/hysteria-core/internal/frag/frag.go new file mode 100644 index 0000000..e7520d9 --- /dev/null +++ b/third_party/hysteria-core/internal/frag/frag.go @@ -0,0 +1,110 @@ +package frag + +import ( + "errors" + + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +// MaxFragments bounds attacker-controlled defragmentation state. A maximum +// Hysteria UDP payload is 4096 bytes, so 16 fragments leaves ample headroom +// above the normal four-to-five fragments while preventing 255-slot abuse. +const MaxFragments = 16 + +// ErrFragmentationLimit reports that a datagram cannot fit within the bounded +// fragment count. Callers must surface this instead of silently sending zero +// fragments. +var ErrFragmentationLimit = errors.New("UDP message exceeds fragmentation limit") + +func FragUDPMessage(m *protocol.UDPMessage, maxSize int) []protocol.UDPMessage { + if m.Size() <= maxSize { + return []protocol.UDPMessage{*m} + } + fullPayload := m.Data + maxPayloadSize := maxSize - m.HeaderSize() + if maxPayloadSize <= 0 { + return nil + } + off := 0 + fragID := uint8(0) + count := (len(fullPayload) + maxPayloadSize - 1) / maxPayloadSize // round up + if count > MaxFragments { + return nil + } + fragCount := uint8(count) + frags := make([]protocol.UDPMessage, fragCount) + for off < len(fullPayload) { + payloadSize := len(fullPayload) - off + if payloadSize > maxPayloadSize { + payloadSize = maxPayloadSize + } + frag := *m + frag.FragID = fragID + frag.FragCount = fragCount + frag.Data = fullPayload[off : off+payloadSize] + frags[fragID] = frag + off += payloadSize + fragID++ + } + return frags +} + +// Defragger handles the defragmentation of UDP messages. +// The current implementation can only handle one packet ID at a time. +// If another packet arrives before a packet has received all fragments +// in their entirety, any previous state is discarded. +type Defragger struct { + pktID uint16 + frags []*protocol.UDPMessage + count uint8 + size int // data size +} + +func (d *Defragger) Feed(m *protocol.UDPMessage) *protocol.UDPMessage { + if m.FragCount <= 1 { + if len(m.Data) > protocol.MaxUDPSize { + return nil + } + return m + } + if m.FragID >= m.FragCount || m.FragCount > MaxFragments { + // wtf is this? + return nil + } + if m.PacketID != d.pktID || m.FragCount != uint8(len(d.frags)) { + // new message, clear previous state + d.pktID = m.PacketID + d.frags = make([]*protocol.UDPMessage, m.FragCount) + d.frags[m.FragID] = m + d.count = 1 + d.size = len(m.Data) + if d.size > protocol.MaxUDPSize { + d.frags = nil + d.count = 0 + d.size = 0 + } + } else if d.frags[m.FragID] == nil { + if d.size+len(m.Data) > protocol.MaxUDPSize { + d.frags = nil + d.count = 0 + d.size = 0 + return nil + } + d.frags[m.FragID] = m + d.count++ + d.size += len(m.Data) + if int(d.count) == len(d.frags) { + // all fragments received, assemble + data := make([]byte, d.size) + off := 0 + for _, frag := range d.frags { + off += copy(data[off:], frag.Data) + } + m.Data = data + m.FragID = 0 + m.FragCount = 1 + return m + } + } + return nil +} diff --git a/third_party/hysteria-core/internal/frag/frag_test.go b/third_party/hysteria-core/internal/frag/frag_test.go new file mode 100644 index 0000000..3f2eb1d --- /dev/null +++ b/third_party/hysteria-core/internal/frag/frag_test.go @@ -0,0 +1,385 @@ +package frag + +import ( + "reflect" + "testing" + + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +func TestFragUDPMessage(t *testing.T) { + type args struct { + m *protocol.UDPMessage + maxSize int + } + tests := []struct { + name string + args args + want []protocol.UDPMessage + }{ + { + "no frag", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 123, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("hello"), + }, + 100, + }, + []protocol.UDPMessage{ + { + SessionID: 123, + PacketID: 123, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("hello"), + }, + }, + }, + { + "2 frags", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 123, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("hello"), + }, + 20, + }, + []protocol.UDPMessage{ + { + SessionID: 123, + PacketID: 123, + FragID: 0, + FragCount: 2, + Addr: "test:123", + Data: []byte("hel"), + }, + { + SessionID: 123, + PacketID: 123, + FragID: 1, + FragCount: 2, + Addr: "test:123", + Data: []byte("lo"), + }, + }, + }, + { + "4 frags", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 123, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("abcdefgh"), + }, + 19, + }, + []protocol.UDPMessage{ + { + SessionID: 123, + PacketID: 123, + FragID: 0, + FragCount: 4, + Addr: "test:123", + Data: []byte("ab"), + }, + { + SessionID: 123, + PacketID: 123, + FragID: 1, + FragCount: 4, + Addr: "test:123", + Data: []byte("cd"), + }, + { + SessionID: 123, + PacketID: 123, + FragID: 2, + FragCount: 4, + Addr: "test:123", + Data: []byte("ef"), + }, + { + SessionID: 123, + PacketID: 123, + FragID: 3, + FragCount: 4, + Addr: "test:123", + Data: []byte("gh"), + }, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := FragUDPMessage(tt.args.m, tt.args.maxSize); !reflect.DeepEqual(got, tt.want) { + t.Errorf("FragUDPMessage() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestDefragger(t *testing.T) { + type args struct { + m *protocol.UDPMessage + } + tests := []struct { + name string + args args + want *protocol.UDPMessage + }{ + { + "no frag", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("hello"), + }, + }, + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("hello"), + }, + }, + { + "frag 0 - 1/2", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 0, + FragCount: 2, + Addr: "test:123", + Data: []byte("hello "), + }, + }, + nil, + }, + { + "frag 0 - 2/2", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 1, + FragCount: 2, + Addr: "test:123", + Data: []byte("moto"), + }, + }, + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("hello moto"), + }, + }, + { + "frag 1 - 1/3", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 0, + FragCount: 3, + Addr: "test:123", + Data: []byte("deco"), + }, + }, + nil, + }, + { + "frag 1 - 2/3", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 1, + FragCount: 3, + Addr: "test:123", + Data: []byte("*"), + }, + }, + nil, + }, + { + "frag 1 - 3/3", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 2, + FragCount: 3, + Addr: "test:123", + Data: []byte("27"), + }, + }, + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 987, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("deco*27"), + }, + }, + { + "frag 2 - 1/2", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 233, + FragID: 1, + FragCount: 2, + Addr: "test:123", + Data: []byte("shinsekai"), + }, + }, + nil, + }, + { + "frag 3 - 2/2", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 244, + FragID: 1, + FragCount: 2, + Addr: "test:123", + Data: []byte("what???"), + }, + }, + nil, + }, + { + "frag 2 - 2/2", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 233, + FragID: 1, + FragCount: 2, + Addr: "test:123", + Data: []byte(" annaijo"), + }, + }, + nil, + }, + { + "invalid id", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 233, + FragID: 88, + FragCount: 2, + Addr: "test:123", + Data: []byte("shinsekai"), + }, + }, + nil, + }, + { + "frag 2 - 1/2 re", + args{ + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 233, + FragID: 0, + FragCount: 2, + Addr: "test:123", + Data: []byte("shinsekai"), + }, + }, + &protocol.UDPMessage{ + SessionID: 123, + PacketID: 233, + FragID: 0, + FragCount: 1, + Addr: "test:123", + Data: []byte("shinsekai annaijo"), + }, + }, + } + + d := &Defragger{} + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := d.Feed(tt.args.m); !reflect.DeepEqual(got, tt.want) { + t.Errorf("Feed() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestDefraggerBoundsAttackerControlledState(t *testing.T) { + d := &Defragger{} + if got := d.Feed(&protocol.UDPMessage{ + SessionID: 1, + PacketID: 1, + FragID: 0, + FragCount: MaxFragments + 1, + Data: []byte("partial"), + }); got != nil { + t.Fatal("excessive fragment count was accepted") + } + if len(d.frags) != 0 { + t.Fatalf("excessive fragment count allocated %d slots", len(d.frags)) + } + + first := make([]byte, protocol.MaxUDPSize-1) + if got := d.Feed(&protocol.UDPMessage{ + SessionID: 2, + PacketID: 2, + FragID: 0, + FragCount: 2, + Data: first, + }); got != nil { + t.Fatal("incomplete message unexpectedly assembled") + } + if got := d.Feed(&protocol.UDPMessage{ + SessionID: 2, + PacketID: 2, + FragID: 1, + FragCount: 2, + Data: []byte("too large"), + }); got != nil { + t.Fatal("oversized reassembled message was accepted") + } + if len(d.frags) != 0 || d.size != 0 || d.count != 0 { + t.Fatalf("oversized state was retained: fragments=%d size=%d count=%d", len(d.frags), d.size, d.count) + } +} + +func TestFragUDPMessageRejectsExcessiveFragmentCount(t *testing.T) { + message := &protocol.UDPMessage{Data: make([]byte, protocol.MaxUDPSize)} + // A deliberately tiny payload budget used to overflow uint8 when the + // fragment count was calculated before validation. + maxSize := message.HeaderSize() + 1 + if got := FragUDPMessage(message, maxSize); got != nil { + t.Fatalf("created %d fragments, want bounded rejection", len(got)) + } +} diff --git a/third_party/hysteria-core/internal/integration_tests/.mockery.yaml b/third_party/hysteria-core/internal/integration_tests/.mockery.yaml new file mode 100644 index 0000000..550a725 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/.mockery.yaml @@ -0,0 +1,29 @@ +with-expecter: true +dir: mocks +outpkg: mocks +packages: + net: + interfaces: + Conn: + config: + mockname: MockConn + github.com/apernet/hysteria/core/v2/server: + interfaces: + Outbound: + config: + mockname: MockOutbound + UDPConn: + config: + mockname: MockUDPConn + Authenticator: + config: + mockname: MockAuthenticator + EventLogger: + config: + mockname: MockEventLogger + TrafficLogger: + config: + mockname: MockTrafficLogger + RequestHook: + config: + mockname: MockRequestHook \ No newline at end of file diff --git a/third_party/hysteria-core/internal/integration_tests/chrome_parrot_test.go b/third_party/hysteria-core/internal/integration_tests/chrome_parrot_test.go new file mode 100644 index 0000000..ea57b0f --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/chrome_parrot_test.go @@ -0,0 +1,70 @@ +package integration_tests + +import ( + "io" + "net" + "testing" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +// TestClientServerChromeParrot runs a real Hysteria client/server pair both with +// the Chrome handshake fingerprint (the default) and with it turned off, proving +// each works end to end through Hysteria's own config plumbing rather than only +// at the quic-go layer. +func TestClientServerChromeParrot(t *testing.T) { + tests := []struct { + name string + quicConfig client.QUICConfig + }{ + {"default", client.QUICConfig{}}, + {"disabled", client.QUICConfig{DisableChromeParrot: true}}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + echoAddr := "127.0.0.1:22444" + echoListener, err := net.Listen("tcp", echoAddr) + assert.NoError(t, err) + echoServer := &tcpEchoServer{Listener: echoListener} + defer echoServer.Close() + go echoServer.Serve() + + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + QUICConfig: test.quicConfig, + }) + assert.NoError(t, err) + defer c.Close() + + conn, err := c.TCP(echoAddr) + assert.NoError(t, err) + defer conn.Close() + + sData := []byte("hello from a chrome-shaped handshake") + _, err = conn.Write(sData) + assert.NoError(t, err) + rData := make([]byte, len(sData)) + _, err = io.ReadFull(conn, rData) + assert.NoError(t, err) + assert.Equal(t, sData, rData) + }) + } +} diff --git a/third_party/hysteria-core/internal/integration_tests/close_test.go b/third_party/hysteria-core/internal/integration_tests/close_test.go new file mode 100644 index 0000000..ac7f84b --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/close_test.go @@ -0,0 +1,252 @@ +package integration_tests + +import ( + "io" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" +) + +// TestClientServerTCPClose tests whether the client/server propagates the close of a connection correctly. +// Closing one side of the connection should close the other side as well. +func TestClientServerTCPClose(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + serverOb := mocks.NewMockOutbound(t) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Outbound: serverOb, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + addr := "hi-and-goodbye:2333" + + // Test close from client side: + // Client creates a connection, writes something, then closes it. + // Server outbound connection should write the same thing, then close. + sobConn := mocks.NewMockConn(t) + sobConnCh := make(chan struct{}) // For close signal only + sobConnChCloseFunc := sync.OnceFunc(func() { close(sobConnCh) }) + sobConn.EXPECT().Read(mock.Anything).RunAndReturn(func(bs []byte) (int, error) { + <-sobConnCh + return 0, io.EOF + }) + sobConn.EXPECT().Write([]byte("happy")).Return(5, nil) + sobConn.EXPECT().Close().RunAndReturn(func() error { + sobConnChCloseFunc() + return nil + }) + serverOb.EXPECT().TCP(addr).Return(sobConn, nil).Once() + conn, err := c.TCP(addr) + assert.NoError(t, err) + _, err = conn.Write([]byte("happy")) + assert.NoError(t, err) + err = conn.Close() + assert.NoError(t, err) + time.Sleep(1 * time.Second) + mock.AssertExpectationsForObjects(t, sobConn, serverOb) + + // Test close from server side: + // Client creates a connection. + // Server outbound connection reads something, then closes. + // Client connection should read the same thing, then close. + sobConn = mocks.NewMockConn(t) + sobConnCh2 := make(chan []byte, 1) + sobConn.EXPECT().Read(mock.Anything).RunAndReturn(func(bs []byte) (int, error) { + d := <-sobConnCh2 + if d == nil { + return 0, io.EOF + } else { + return copy(bs, d), nil + } + }) + sobConn.EXPECT().Close().Return(nil) + serverOb.EXPECT().TCP(addr).Return(sobConn, nil).Once() + conn, err = c.TCP(addr) + assert.NoError(t, err) + sobConnCh2 <- []byte("happy") + close(sobConnCh2) + bs, err := io.ReadAll(conn) + assert.NoError(t, err) + assert.Equal(t, "happy", string(bs)) +} + +// TestClientServerUDPIdleTimeout tests whether the server's UDP idle timeout works correctly. +func TestClientServerUDPIdleTimeout(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + serverOb := mocks.NewMockOutbound(t) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + eventLogger := mocks.NewMockEventLogger(t) + eventLogger.EXPECT().Connect(mock.Anything, "nobody", mock.Anything).Once() + eventLogger.EXPECT().Disconnect(mock.Anything, "nobody", mock.Anything).Maybe() // Depends on the timing, don't care + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Outbound: serverOb, + UDPIdleTimeout: 2 * time.Second, + Authenticator: auth, + EventLogger: eventLogger, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + addr := "spy.x.family:2023" + + // On the client side, create a UDP session and send a packet every 1 second, + // 4 packets in total. The server should have one UDP session and receive all + // 4 packets. Then the UDP connection on the server side will receive a packet + // every 1 second, 4 packets in total. The client session should receive all + // 4 packets. Then the session will be idle for 3 seconds - should be enough + // to trigger the server's UDP idle timeout. + sobConn := mocks.NewMockUDPConn(t) + sobConnCh := make(chan []byte, 1) + sobConnChCloseFunc := sync.OnceFunc(func() { close(sobConnCh) }) + sobConn.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(bs []byte) (int, string, error) { + d := <-sobConnCh + if d == nil { + return 0, "", io.EOF + } else { + return copy(bs, d), addr, nil + } + }) + sobConn.EXPECT().WriteTo([]byte("happy"), addr).Return(5, nil).Times(4) + serverOb.EXPECT().UDP(addr).Return(sobConn, nil).Once() + eventLogger.EXPECT().UDPRequest(mock.Anything, mock.Anything, uint32(1), addr).Once() + cu, err := c.UDP() + assert.NoError(t, err) + // Client sends 4 packets + for i := 0; i < 4; i++ { + err = cu.Send([]byte("happy"), addr) + assert.NoError(t, err) + time.Sleep(1 * time.Second) + } + // Client receives 4 packets + go func() { + for i := 0; i < 4; i++ { + sobConnCh <- []byte("sad") + time.Sleep(1 * time.Second) + } + }() + for i := 0; i < 4; i++ { + bs, rAddr, err := cu.Receive() + assert.NoError(t, err) + assert.Equal(t, "sad", string(bs)) + assert.Equal(t, addr, rAddr) + } + // Now we wait for 3 seconds, the server should close the UDP session. + sobConn.EXPECT().Close().RunAndReturn(func() error { + sobConnChCloseFunc() + return nil + }) + eventLogger.EXPECT().UDPError(mock.Anything, mock.Anything, uint32(1), nil).Once() + time.Sleep(3 * time.Second) +} + +// TestClientServerClientShutdown tests whether the server can handle the client's shutdown correctly. +func TestClientServerClientShutdown(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + eventLogger := mocks.NewMockEventLogger(t) + eventLogger.EXPECT().Connect(mock.Anything, "nobody", mock.Anything).Once() + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + EventLogger: eventLogger, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + + // Close the client - expect disconnect event on the server side. + // Since client.Close() sends HTTP3 ErrCodeNoError, the error should be nil. + eventLogger.EXPECT().Disconnect(mock.Anything, "nobody", nil).Once() + _ = c.Close() + time.Sleep(1 * time.Second) +} + +// TestClientServerServerShutdown tests whether the client can handle the server's shutdown correctly. +func TestClientServerServerShutdown(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + go s.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + QUICConfig: client.QUICConfig{ + MaxIdleTimeout: 4 * time.Second, + }, + }) + assert.NoError(t, err) + + // Close the server - expect the client to return ClosedError for both TCP & UDP calls. + _ = s.Close() + + _, err = c.TCP("whatever") + _, ok := err.(errors.ClosedError) + assert.True(t, ok) + + time.Sleep(1 * time.Second) // Allow some time for the error to be propagated to the UDP session manager + + _, err = c.UDP() + _, ok = err.(errors.ClosedError) + assert.True(t, ok) + + assert.NoError(t, c.Close()) +} diff --git a/third_party/hysteria-core/internal/integration_tests/hook_test.go b/third_party/hysteria-core/internal/integration_tests/hook_test.go new file mode 100644 index 0000000..db95995 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/hook_test.go @@ -0,0 +1,146 @@ +package integration_tests + +import ( + "io" + "net" + "testing" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +func TestClientServerHookTCP(t *testing.T) { + fakeEchoAddr := "hahanope:6666" + realEchoAddr := "127.0.0.1:22333" + + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + hook := mocks.NewMockRequestHook(t) + hook.EXPECT().Check(false, fakeEchoAddr).Return(true).Once() + hook.EXPECT().TCP(mock.Anything, mock.Anything).RunAndReturn(func(stream server.HyStream, s *string) ([]byte, error) { + assert.Equal(t, fakeEchoAddr, *s) + // Change the address + *s = realEchoAddr + // Read the first 5 bytes and replace them with "byeee" + data := make([]byte, 5) + _, err := io.ReadFull(stream, data) + if err != nil { + return nil, err + } + assert.Equal(t, []byte("hello"), data) + return []byte("byeee"), nil + }).Once() + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + RequestHook: hook, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create TCP echo server + echoListener, err := net.Listen("tcp", realEchoAddr) + assert.NoError(t, err) + echoServer := &tcpEchoServer{Listener: echoListener} + defer echoServer.Close() + go echoServer.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + // Dial TCP + conn, err := c.TCP(fakeEchoAddr) + assert.NoError(t, err) + defer conn.Close() + + // Send and receive data + sData := []byte("hello world") + _, err = conn.Write(sData) + assert.NoError(t, err) + rData := make([]byte, len(sData)) + _, err = io.ReadFull(conn, rData) + assert.NoError(t, err) + assert.Equal(t, []byte("byeee world"), rData) +} + +func TestClientServerHookUDP(t *testing.T) { + fakeEchoAddr := "hahanope:6666" + realEchoAddr := "127.0.0.1:22333" + + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + hook := mocks.NewMockRequestHook(t) + hook.EXPECT().Check(true, fakeEchoAddr).Return(true).Once() + hook.EXPECT().UDP(mock.Anything, mock.Anything).RunAndReturn(func(bytes []byte, s *string) error { + assert.Equal(t, fakeEchoAddr, *s) + assert.Equal(t, []byte("hello world"), bytes) + // Change the address + *s = realEchoAddr + return nil + }).Once() + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + RequestHook: hook, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create UDP echo server + echoConn, err := net.ListenPacket("udp", realEchoAddr) + assert.NoError(t, err) + echoServer := &udpEchoServer{Conn: echoConn} + defer echoServer.Close() + go echoServer.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + // Listen UDP + conn, err := c.UDP() + assert.NoError(t, err) + defer conn.Close() + + // Send and receive data + sData := []byte("hello world") + err = conn.Send(sData, fakeEchoAddr) + assert.NoError(t, err) + rData, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, sData, rData) + // Hook address change is transparent, + // the client should still see the fake echo address it sent packets to + assert.Equal(t, fakeEchoAddr, rAddr) + + // Subsequent packets should also be sent to the real echo server + sData = []byte("never stop fighting") + err = conn.Send(sData, fakeEchoAddr) + assert.NoError(t, err) + rData, rAddr, err = conn.Receive() + assert.NoError(t, err) + assert.Equal(t, sData, rData) + assert.Equal(t, fakeEchoAddr, rAddr) +} diff --git a/third_party/hysteria-core/internal/integration_tests/masq_test.go b/third_party/hysteria-core/internal/integration_tests/masq_test.go new file mode 100644 index 0000000..584e2f1 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/masq_test.go @@ -0,0 +1,93 @@ +package integration_tests + +import ( + "context" + "crypto/tls" + "net" + "net/http" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/internal/protocol" + "github.com/apernet/hysteria/core/v2/server" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" +) + +// TestServerMasquerade is a test to ensure that the server behaves as a normal +// HTTP/3 server when dealing with an unauthenticated client. This is mainly to +// confirm that the server does not expose itself to active probing. +func TestServerMasquerade(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, "", uint64(0)).Return(false, "").Once() + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // QUIC connection & RoundTripper + var conn *quic.Conn + rt := &http3.Transport{ + TLSClientConfig: &tls.Config{ + InsecureSkipVerify: true, + }, + Dial: func(ctx context.Context, _ string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + qc, err := quic.DialAddrEarly(ctx, udpAddr.String(), tlsCfg, cfg) + if err != nil { + return nil, err + } + conn = qc + return qc, nil + }, + } + defer rt.Close() // This will close the QUIC connection + + // Send the bogus request + // We expect 404 (from the default handler) + req := &http.Request{ + Method: http.MethodPost, + URL: &url.URL{ + Scheme: "https", + Host: protocol.URLHost, + Path: protocol.URLPath, + }, + Header: make(http.Header), + } + resp, err := rt.RoundTrip(req) + assert.NoError(t, err) + assert.Equal(t, http.StatusNotFound, resp.StatusCode) + for k := range resp.Header { + // Make sure no strange headers are sent by the server + assert.NotContains(t, k, "Hysteria") + } + + buf := make([]byte, 1024) + + // We send a TCP request anyway, see if we get a response + tcpStream, err := conn.OpenStream() + assert.NoError(t, err) + defer tcpStream.Close() + err = protocol.WriteTCPRequest(tcpStream, "www.google.com:443") + assert.NoError(t, err) + + // We should receive nothing + _ = tcpStream.SetReadDeadline(time.Now().Add(2 * time.Second)) + n, err := tcpStream.Read(buf) + assert.Equal(t, 0, n) + nErr, ok := err.(net.Error) + assert.True(t, ok) + assert.True(t, nErr.Timeout()) +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_Authenticator.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_Authenticator.go new file mode 100644 index 0000000..018b499 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_Authenticator.go @@ -0,0 +1,94 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import ( + net "net" + + mock "github.com/stretchr/testify/mock" +) + +// MockAuthenticator is an autogenerated mock type for the Authenticator type +type MockAuthenticator struct { + mock.Mock +} + +type MockAuthenticator_Expecter struct { + mock *mock.Mock +} + +func (_m *MockAuthenticator) EXPECT() *MockAuthenticator_Expecter { + return &MockAuthenticator_Expecter{mock: &_m.Mock} +} + +// Authenticate provides a mock function with given fields: addr, auth, tx +func (_m *MockAuthenticator) Authenticate(addr net.Addr, auth string, tx uint64) (bool, string) { + ret := _m.Called(addr, auth, tx) + + if len(ret) == 0 { + panic("no return value specified for Authenticate") + } + + var r0 bool + var r1 string + if rf, ok := ret.Get(0).(func(net.Addr, string, uint64) (bool, string)); ok { + return rf(addr, auth, tx) + } + if rf, ok := ret.Get(0).(func(net.Addr, string, uint64) bool); ok { + r0 = rf(addr, auth, tx) + } else { + r0 = ret.Get(0).(bool) + } + + if rf, ok := ret.Get(1).(func(net.Addr, string, uint64) string); ok { + r1 = rf(addr, auth, tx) + } else { + r1 = ret.Get(1).(string) + } + + return r0, r1 +} + +// MockAuthenticator_Authenticate_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Authenticate' +type MockAuthenticator_Authenticate_Call struct { + *mock.Call +} + +// Authenticate is a helper method to define mock.On call +// - addr net.Addr +// - auth string +// - tx uint64 +func (_e *MockAuthenticator_Expecter) Authenticate(addr interface{}, auth interface{}, tx interface{}) *MockAuthenticator_Authenticate_Call { + return &MockAuthenticator_Authenticate_Call{Call: _e.mock.On("Authenticate", addr, auth, tx)} +} + +func (_c *MockAuthenticator_Authenticate_Call) Run(run func(addr net.Addr, auth string, tx uint64)) *MockAuthenticator_Authenticate_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(uint64)) + }) + return _c +} + +func (_c *MockAuthenticator_Authenticate_Call) Return(ok bool, id string) *MockAuthenticator_Authenticate_Call { + _c.Call.Return(ok, id) + return _c +} + +func (_c *MockAuthenticator_Authenticate_Call) RunAndReturn(run func(net.Addr, string, uint64) (bool, string)) *MockAuthenticator_Authenticate_Call { + _c.Call.Return(run) + return _c +} + +// NewMockAuthenticator creates a new instance of MockAuthenticator. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockAuthenticator(t interface { + mock.TestingT + Cleanup(func()) +}) *MockAuthenticator { + mock := &MockAuthenticator{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_Conn.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_Conn.go new file mode 100644 index 0000000..d068033 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_Conn.go @@ -0,0 +1,426 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import ( + net "net" + time "time" + + mock "github.com/stretchr/testify/mock" +) + +// MockConn is an autogenerated mock type for the Conn type +type MockConn struct { + mock.Mock +} + +type MockConn_Expecter struct { + mock *mock.Mock +} + +func (_m *MockConn) EXPECT() *MockConn_Expecter { + return &MockConn_Expecter{mock: &_m.Mock} +} + +// Close provides a mock function with no fields +func (_m *MockConn) Close() error { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for Close") + } + + var r0 error + if rf, ok := ret.Get(0).(func() error); ok { + r0 = rf() + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockConn_Close_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Close' +type MockConn_Close_Call struct { + *mock.Call +} + +// Close is a helper method to define mock.On call +func (_e *MockConn_Expecter) Close() *MockConn_Close_Call { + return &MockConn_Close_Call{Call: _e.mock.On("Close")} +} + +func (_c *MockConn_Close_Call) Run(run func()) *MockConn_Close_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *MockConn_Close_Call) Return(_a0 error) *MockConn_Close_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockConn_Close_Call) RunAndReturn(run func() error) *MockConn_Close_Call { + _c.Call.Return(run) + return _c +} + +// LocalAddr provides a mock function with no fields +func (_m *MockConn) LocalAddr() net.Addr { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for LocalAddr") + } + + var r0 net.Addr + if rf, ok := ret.Get(0).(func() net.Addr); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(net.Addr) + } + } + + return r0 +} + +// MockConn_LocalAddr_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'LocalAddr' +type MockConn_LocalAddr_Call struct { + *mock.Call +} + +// LocalAddr is a helper method to define mock.On call +func (_e *MockConn_Expecter) LocalAddr() *MockConn_LocalAddr_Call { + return &MockConn_LocalAddr_Call{Call: _e.mock.On("LocalAddr")} +} + +func (_c *MockConn_LocalAddr_Call) Run(run func()) *MockConn_LocalAddr_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *MockConn_LocalAddr_Call) Return(_a0 net.Addr) *MockConn_LocalAddr_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockConn_LocalAddr_Call) RunAndReturn(run func() net.Addr) *MockConn_LocalAddr_Call { + _c.Call.Return(run) + return _c +} + +// Read provides a mock function with given fields: b +func (_m *MockConn) Read(b []byte) (int, error) { + ret := _m.Called(b) + + if len(ret) == 0 { + panic("no return value specified for Read") + } + + var r0 int + var r1 error + if rf, ok := ret.Get(0).(func([]byte) (int, error)); ok { + return rf(b) + } + if rf, ok := ret.Get(0).(func([]byte) int); ok { + r0 = rf(b) + } else { + r0 = ret.Get(0).(int) + } + + if rf, ok := ret.Get(1).(func([]byte) error); ok { + r1 = rf(b) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MockConn_Read_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Read' +type MockConn_Read_Call struct { + *mock.Call +} + +// Read is a helper method to define mock.On call +// - b []byte +func (_e *MockConn_Expecter) Read(b interface{}) *MockConn_Read_Call { + return &MockConn_Read_Call{Call: _e.mock.On("Read", b)} +} + +func (_c *MockConn_Read_Call) Run(run func(b []byte)) *MockConn_Read_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte)) + }) + return _c +} + +func (_c *MockConn_Read_Call) Return(n int, err error) *MockConn_Read_Call { + _c.Call.Return(n, err) + return _c +} + +func (_c *MockConn_Read_Call) RunAndReturn(run func([]byte) (int, error)) *MockConn_Read_Call { + _c.Call.Return(run) + return _c +} + +// RemoteAddr provides a mock function with no fields +func (_m *MockConn) RemoteAddr() net.Addr { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for RemoteAddr") + } + + var r0 net.Addr + if rf, ok := ret.Get(0).(func() net.Addr); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(net.Addr) + } + } + + return r0 +} + +// MockConn_RemoteAddr_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoteAddr' +type MockConn_RemoteAddr_Call struct { + *mock.Call +} + +// RemoteAddr is a helper method to define mock.On call +func (_e *MockConn_Expecter) RemoteAddr() *MockConn_RemoteAddr_Call { + return &MockConn_RemoteAddr_Call{Call: _e.mock.On("RemoteAddr")} +} + +func (_c *MockConn_RemoteAddr_Call) Run(run func()) *MockConn_RemoteAddr_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *MockConn_RemoteAddr_Call) Return(_a0 net.Addr) *MockConn_RemoteAddr_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockConn_RemoteAddr_Call) RunAndReturn(run func() net.Addr) *MockConn_RemoteAddr_Call { + _c.Call.Return(run) + return _c +} + +// SetDeadline provides a mock function with given fields: t +func (_m *MockConn) SetDeadline(t time.Time) error { + ret := _m.Called(t) + + if len(ret) == 0 { + panic("no return value specified for SetDeadline") + } + + var r0 error + if rf, ok := ret.Get(0).(func(time.Time) error); ok { + r0 = rf(t) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockConn_SetDeadline_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetDeadline' +type MockConn_SetDeadline_Call struct { + *mock.Call +} + +// SetDeadline is a helper method to define mock.On call +// - t time.Time +func (_e *MockConn_Expecter) SetDeadline(t interface{}) *MockConn_SetDeadline_Call { + return &MockConn_SetDeadline_Call{Call: _e.mock.On("SetDeadline", t)} +} + +func (_c *MockConn_SetDeadline_Call) Run(run func(t time.Time)) *MockConn_SetDeadline_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(time.Time)) + }) + return _c +} + +func (_c *MockConn_SetDeadline_Call) Return(_a0 error) *MockConn_SetDeadline_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockConn_SetDeadline_Call) RunAndReturn(run func(time.Time) error) *MockConn_SetDeadline_Call { + _c.Call.Return(run) + return _c +} + +// SetReadDeadline provides a mock function with given fields: t +func (_m *MockConn) SetReadDeadline(t time.Time) error { + ret := _m.Called(t) + + if len(ret) == 0 { + panic("no return value specified for SetReadDeadline") + } + + var r0 error + if rf, ok := ret.Get(0).(func(time.Time) error); ok { + r0 = rf(t) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockConn_SetReadDeadline_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetReadDeadline' +type MockConn_SetReadDeadline_Call struct { + *mock.Call +} + +// SetReadDeadline is a helper method to define mock.On call +// - t time.Time +func (_e *MockConn_Expecter) SetReadDeadline(t interface{}) *MockConn_SetReadDeadline_Call { + return &MockConn_SetReadDeadline_Call{Call: _e.mock.On("SetReadDeadline", t)} +} + +func (_c *MockConn_SetReadDeadline_Call) Run(run func(t time.Time)) *MockConn_SetReadDeadline_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(time.Time)) + }) + return _c +} + +func (_c *MockConn_SetReadDeadline_Call) Return(_a0 error) *MockConn_SetReadDeadline_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockConn_SetReadDeadline_Call) RunAndReturn(run func(time.Time) error) *MockConn_SetReadDeadline_Call { + _c.Call.Return(run) + return _c +} + +// SetWriteDeadline provides a mock function with given fields: t +func (_m *MockConn) SetWriteDeadline(t time.Time) error { + ret := _m.Called(t) + + if len(ret) == 0 { + panic("no return value specified for SetWriteDeadline") + } + + var r0 error + if rf, ok := ret.Get(0).(func(time.Time) error); ok { + r0 = rf(t) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockConn_SetWriteDeadline_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetWriteDeadline' +type MockConn_SetWriteDeadline_Call struct { + *mock.Call +} + +// SetWriteDeadline is a helper method to define mock.On call +// - t time.Time +func (_e *MockConn_Expecter) SetWriteDeadline(t interface{}) *MockConn_SetWriteDeadline_Call { + return &MockConn_SetWriteDeadline_Call{Call: _e.mock.On("SetWriteDeadline", t)} +} + +func (_c *MockConn_SetWriteDeadline_Call) Run(run func(t time.Time)) *MockConn_SetWriteDeadline_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(time.Time)) + }) + return _c +} + +func (_c *MockConn_SetWriteDeadline_Call) Return(_a0 error) *MockConn_SetWriteDeadline_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockConn_SetWriteDeadline_Call) RunAndReturn(run func(time.Time) error) *MockConn_SetWriteDeadline_Call { + _c.Call.Return(run) + return _c +} + +// Write provides a mock function with given fields: b +func (_m *MockConn) Write(b []byte) (int, error) { + ret := _m.Called(b) + + if len(ret) == 0 { + panic("no return value specified for Write") + } + + var r0 int + var r1 error + if rf, ok := ret.Get(0).(func([]byte) (int, error)); ok { + return rf(b) + } + if rf, ok := ret.Get(0).(func([]byte) int); ok { + r0 = rf(b) + } else { + r0 = ret.Get(0).(int) + } + + if rf, ok := ret.Get(1).(func([]byte) error); ok { + r1 = rf(b) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MockConn_Write_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Write' +type MockConn_Write_Call struct { + *mock.Call +} + +// Write is a helper method to define mock.On call +// - b []byte +func (_e *MockConn_Expecter) Write(b interface{}) *MockConn_Write_Call { + return &MockConn_Write_Call{Call: _e.mock.On("Write", b)} +} + +func (_c *MockConn_Write_Call) Run(run func(b []byte)) *MockConn_Write_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte)) + }) + return _c +} + +func (_c *MockConn_Write_Call) Return(n int, err error) *MockConn_Write_Call { + _c.Call.Return(n, err) + return _c +} + +func (_c *MockConn_Write_Call) RunAndReturn(run func([]byte) (int, error)) *MockConn_Write_Call { + _c.Call.Return(run) + return _c +} + +// NewMockConn creates a new instance of MockConn. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockConn(t interface { + mock.TestingT + Cleanup(func()) +}) *MockConn { + mock := &MockConn{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_EventLogger.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_EventLogger.go new file mode 100644 index 0000000..14f2175 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_EventLogger.go @@ -0,0 +1,249 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import ( + net "net" + + mock "github.com/stretchr/testify/mock" +) + +// MockEventLogger is an autogenerated mock type for the EventLogger type +type MockEventLogger struct { + mock.Mock +} + +type MockEventLogger_Expecter struct { + mock *mock.Mock +} + +func (_m *MockEventLogger) EXPECT() *MockEventLogger_Expecter { + return &MockEventLogger_Expecter{mock: &_m.Mock} +} + +// Connect provides a mock function with given fields: addr, id, tx +func (_m *MockEventLogger) Connect(addr net.Addr, id string, tx uint64) { + _m.Called(addr, id, tx) +} + +// MockEventLogger_Connect_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Connect' +type MockEventLogger_Connect_Call struct { + *mock.Call +} + +// Connect is a helper method to define mock.On call +// - addr net.Addr +// - id string +// - tx uint64 +func (_e *MockEventLogger_Expecter) Connect(addr interface{}, id interface{}, tx interface{}) *MockEventLogger_Connect_Call { + return &MockEventLogger_Connect_Call{Call: _e.mock.On("Connect", addr, id, tx)} +} + +func (_c *MockEventLogger_Connect_Call) Run(run func(addr net.Addr, id string, tx uint64)) *MockEventLogger_Connect_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(uint64)) + }) + return _c +} + +func (_c *MockEventLogger_Connect_Call) Return() *MockEventLogger_Connect_Call { + _c.Call.Return() + return _c +} + +func (_c *MockEventLogger_Connect_Call) RunAndReturn(run func(net.Addr, string, uint64)) *MockEventLogger_Connect_Call { + _c.Run(run) + return _c +} + +// Disconnect provides a mock function with given fields: addr, id, err +func (_m *MockEventLogger) Disconnect(addr net.Addr, id string, err error) { + _m.Called(addr, id, err) +} + +// MockEventLogger_Disconnect_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Disconnect' +type MockEventLogger_Disconnect_Call struct { + *mock.Call +} + +// Disconnect is a helper method to define mock.On call +// - addr net.Addr +// - id string +// - err error +func (_e *MockEventLogger_Expecter) Disconnect(addr interface{}, id interface{}, err interface{}) *MockEventLogger_Disconnect_Call { + return &MockEventLogger_Disconnect_Call{Call: _e.mock.On("Disconnect", addr, id, err)} +} + +func (_c *MockEventLogger_Disconnect_Call) Run(run func(addr net.Addr, id string, err error)) *MockEventLogger_Disconnect_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(error)) + }) + return _c +} + +func (_c *MockEventLogger_Disconnect_Call) Return() *MockEventLogger_Disconnect_Call { + _c.Call.Return() + return _c +} + +func (_c *MockEventLogger_Disconnect_Call) RunAndReturn(run func(net.Addr, string, error)) *MockEventLogger_Disconnect_Call { + _c.Run(run) + return _c +} + +// TCPError provides a mock function with given fields: addr, id, reqAddr, err +func (_m *MockEventLogger) TCPError(addr net.Addr, id string, reqAddr string, err error) { + _m.Called(addr, id, reqAddr, err) +} + +// MockEventLogger_TCPError_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'TCPError' +type MockEventLogger_TCPError_Call struct { + *mock.Call +} + +// TCPError is a helper method to define mock.On call +// - addr net.Addr +// - id string +// - reqAddr string +// - err error +func (_e *MockEventLogger_Expecter) TCPError(addr interface{}, id interface{}, reqAddr interface{}, err interface{}) *MockEventLogger_TCPError_Call { + return &MockEventLogger_TCPError_Call{Call: _e.mock.On("TCPError", addr, id, reqAddr, err)} +} + +func (_c *MockEventLogger_TCPError_Call) Run(run func(addr net.Addr, id string, reqAddr string, err error)) *MockEventLogger_TCPError_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(string), args[3].(error)) + }) + return _c +} + +func (_c *MockEventLogger_TCPError_Call) Return() *MockEventLogger_TCPError_Call { + _c.Call.Return() + return _c +} + +func (_c *MockEventLogger_TCPError_Call) RunAndReturn(run func(net.Addr, string, string, error)) *MockEventLogger_TCPError_Call { + _c.Run(run) + return _c +} + +// TCPRequest provides a mock function with given fields: addr, id, reqAddr +func (_m *MockEventLogger) TCPRequest(addr net.Addr, id string, reqAddr string) { + _m.Called(addr, id, reqAddr) +} + +// MockEventLogger_TCPRequest_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'TCPRequest' +type MockEventLogger_TCPRequest_Call struct { + *mock.Call +} + +// TCPRequest is a helper method to define mock.On call +// - addr net.Addr +// - id string +// - reqAddr string +func (_e *MockEventLogger_Expecter) TCPRequest(addr interface{}, id interface{}, reqAddr interface{}) *MockEventLogger_TCPRequest_Call { + return &MockEventLogger_TCPRequest_Call{Call: _e.mock.On("TCPRequest", addr, id, reqAddr)} +} + +func (_c *MockEventLogger_TCPRequest_Call) Run(run func(addr net.Addr, id string, reqAddr string)) *MockEventLogger_TCPRequest_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(string)) + }) + return _c +} + +func (_c *MockEventLogger_TCPRequest_Call) Return() *MockEventLogger_TCPRequest_Call { + _c.Call.Return() + return _c +} + +func (_c *MockEventLogger_TCPRequest_Call) RunAndReturn(run func(net.Addr, string, string)) *MockEventLogger_TCPRequest_Call { + _c.Run(run) + return _c +} + +// UDPError provides a mock function with given fields: addr, id, sessionID, err +func (_m *MockEventLogger) UDPError(addr net.Addr, id string, sessionID uint32, err error) { + _m.Called(addr, id, sessionID, err) +} + +// MockEventLogger_UDPError_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UDPError' +type MockEventLogger_UDPError_Call struct { + *mock.Call +} + +// UDPError is a helper method to define mock.On call +// - addr net.Addr +// - id string +// - sessionID uint32 +// - err error +func (_e *MockEventLogger_Expecter) UDPError(addr interface{}, id interface{}, sessionID interface{}, err interface{}) *MockEventLogger_UDPError_Call { + return &MockEventLogger_UDPError_Call{Call: _e.mock.On("UDPError", addr, id, sessionID, err)} +} + +func (_c *MockEventLogger_UDPError_Call) Run(run func(addr net.Addr, id string, sessionID uint32, err error)) *MockEventLogger_UDPError_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(uint32), args[3].(error)) + }) + return _c +} + +func (_c *MockEventLogger_UDPError_Call) Return() *MockEventLogger_UDPError_Call { + _c.Call.Return() + return _c +} + +func (_c *MockEventLogger_UDPError_Call) RunAndReturn(run func(net.Addr, string, uint32, error)) *MockEventLogger_UDPError_Call { + _c.Run(run) + return _c +} + +// UDPRequest provides a mock function with given fields: addr, id, sessionID, reqAddr +func (_m *MockEventLogger) UDPRequest(addr net.Addr, id string, sessionID uint32, reqAddr string) { + _m.Called(addr, id, sessionID, reqAddr) +} + +// MockEventLogger_UDPRequest_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UDPRequest' +type MockEventLogger_UDPRequest_Call struct { + *mock.Call +} + +// UDPRequest is a helper method to define mock.On call +// - addr net.Addr +// - id string +// - sessionID uint32 +// - reqAddr string +func (_e *MockEventLogger_Expecter) UDPRequest(addr interface{}, id interface{}, sessionID interface{}, reqAddr interface{}) *MockEventLogger_UDPRequest_Call { + return &MockEventLogger_UDPRequest_Call{Call: _e.mock.On("UDPRequest", addr, id, sessionID, reqAddr)} +} + +func (_c *MockEventLogger_UDPRequest_Call) Run(run func(addr net.Addr, id string, sessionID uint32, reqAddr string)) *MockEventLogger_UDPRequest_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(net.Addr), args[1].(string), args[2].(uint32), args[3].(string)) + }) + return _c +} + +func (_c *MockEventLogger_UDPRequest_Call) Return() *MockEventLogger_UDPRequest_Call { + _c.Call.Return() + return _c +} + +func (_c *MockEventLogger_UDPRequest_Call) RunAndReturn(run func(net.Addr, string, uint32, string)) *MockEventLogger_UDPRequest_Call { + _c.Run(run) + return _c +} + +// NewMockEventLogger creates a new instance of MockEventLogger. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockEventLogger(t interface { + mock.TestingT + Cleanup(func()) +}) *MockEventLogger { + mock := &MockEventLogger{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_Outbound.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_Outbound.go new file mode 100644 index 0000000..6fda640 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_Outbound.go @@ -0,0 +1,199 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import ( + net "net" + + server "github.com/apernet/hysteria/core/v2/server" + mock "github.com/stretchr/testify/mock" +) + +// MockOutbound is an autogenerated mock type for the Outbound type +type MockOutbound struct { + mock.Mock +} + +type MockOutbound_Expecter struct { + mock *mock.Mock +} + +func (_m *MockOutbound) EXPECT() *MockOutbound_Expecter { + return &MockOutbound_Expecter{mock: &_m.Mock} +} + +// CheckUDP provides a mock function with given fields: reqAddr +func (_m *MockOutbound) CheckUDP(reqAddr string) error { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for CheckUDP") + } + + var r0 error + if rf, ok := ret.Get(0).(func(string) error); ok { + r0 = rf(reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockOutbound_CheckUDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckUDP' +type MockOutbound_CheckUDP_Call struct { + *mock.Call +} + +// CheckUDP is a helper method to define mock.On call +// - reqAddr string +func (_e *MockOutbound_Expecter) CheckUDP(reqAddr interface{}) *MockOutbound_CheckUDP_Call { + return &MockOutbound_CheckUDP_Call{Call: _e.mock.On("CheckUDP", reqAddr)} +} + +func (_c *MockOutbound_CheckUDP_Call) Run(run func(reqAddr string)) *MockOutbound_CheckUDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *MockOutbound_CheckUDP_Call) Return(_a0 error) *MockOutbound_CheckUDP_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockOutbound_CheckUDP_Call) RunAndReturn(run func(string) error) *MockOutbound_CheckUDP_Call { + _c.Call.Return(run) + return _c +} + +// TCP provides a mock function with given fields: reqAddr +func (_m *MockOutbound) TCP(reqAddr string) (net.Conn, error) { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for TCP") + } + + var r0 net.Conn + var r1 error + if rf, ok := ret.Get(0).(func(string) (net.Conn, error)); ok { + return rf(reqAddr) + } + if rf, ok := ret.Get(0).(func(string) net.Conn); ok { + r0 = rf(reqAddr) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(net.Conn) + } + } + + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(reqAddr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MockOutbound_TCP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'TCP' +type MockOutbound_TCP_Call struct { + *mock.Call +} + +// TCP is a helper method to define mock.On call +// - reqAddr string +func (_e *MockOutbound_Expecter) TCP(reqAddr interface{}) *MockOutbound_TCP_Call { + return &MockOutbound_TCP_Call{Call: _e.mock.On("TCP", reqAddr)} +} + +func (_c *MockOutbound_TCP_Call) Run(run func(reqAddr string)) *MockOutbound_TCP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *MockOutbound_TCP_Call) Return(_a0 net.Conn, _a1 error) *MockOutbound_TCP_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *MockOutbound_TCP_Call) RunAndReturn(run func(string) (net.Conn, error)) *MockOutbound_TCP_Call { + _c.Call.Return(run) + return _c +} + +// UDP provides a mock function with given fields: reqAddr +func (_m *MockOutbound) UDP(reqAddr string) (server.UDPConn, error) { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for UDP") + } + + var r0 server.UDPConn + var r1 error + if rf, ok := ret.Get(0).(func(string) (server.UDPConn, error)); ok { + return rf(reqAddr) + } + if rf, ok := ret.Get(0).(func(string) server.UDPConn); ok { + r0 = rf(reqAddr) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(server.UDPConn) + } + } + + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(reqAddr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MockOutbound_UDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UDP' +type MockOutbound_UDP_Call struct { + *mock.Call +} + +// UDP is a helper method to define mock.On call +// - reqAddr string +func (_e *MockOutbound_Expecter) UDP(reqAddr interface{}) *MockOutbound_UDP_Call { + return &MockOutbound_UDP_Call{Call: _e.mock.On("UDP", reqAddr)} +} + +func (_c *MockOutbound_UDP_Call) Run(run func(reqAddr string)) *MockOutbound_UDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *MockOutbound_UDP_Call) Return(_a0 server.UDPConn, _a1 error) *MockOutbound_UDP_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *MockOutbound_UDP_Call) RunAndReturn(run func(string) (server.UDPConn, error)) *MockOutbound_UDP_Call { + _c.Call.Return(run) + return _c +} + +// NewMockOutbound creates a new instance of MockOutbound. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockOutbound(t interface { + mock.TestingT + Cleanup(func()) +}) *MockOutbound { + mock := &MockOutbound{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_RequestHook.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_RequestHook.go new file mode 100644 index 0000000..49e8c6c --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_RequestHook.go @@ -0,0 +1,188 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import ( + server "github.com/apernet/hysteria/core/v2/server" + mock "github.com/stretchr/testify/mock" +) + +// MockRequestHook is an autogenerated mock type for the RequestHook type +type MockRequestHook struct { + mock.Mock +} + +type MockRequestHook_Expecter struct { + mock *mock.Mock +} + +func (_m *MockRequestHook) EXPECT() *MockRequestHook_Expecter { + return &MockRequestHook_Expecter{mock: &_m.Mock} +} + +// Check provides a mock function with given fields: isUDP, reqAddr +func (_m *MockRequestHook) Check(isUDP bool, reqAddr string) bool { + ret := _m.Called(isUDP, reqAddr) + + if len(ret) == 0 { + panic("no return value specified for Check") + } + + var r0 bool + if rf, ok := ret.Get(0).(func(bool, string) bool); ok { + r0 = rf(isUDP, reqAddr) + } else { + r0 = ret.Get(0).(bool) + } + + return r0 +} + +// MockRequestHook_Check_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Check' +type MockRequestHook_Check_Call struct { + *mock.Call +} + +// Check is a helper method to define mock.On call +// - isUDP bool +// - reqAddr string +func (_e *MockRequestHook_Expecter) Check(isUDP interface{}, reqAddr interface{}) *MockRequestHook_Check_Call { + return &MockRequestHook_Check_Call{Call: _e.mock.On("Check", isUDP, reqAddr)} +} + +func (_c *MockRequestHook_Check_Call) Run(run func(isUDP bool, reqAddr string)) *MockRequestHook_Check_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(bool), args[1].(string)) + }) + return _c +} + +func (_c *MockRequestHook_Check_Call) Return(_a0 bool) *MockRequestHook_Check_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockRequestHook_Check_Call) RunAndReturn(run func(bool, string) bool) *MockRequestHook_Check_Call { + _c.Call.Return(run) + return _c +} + +// TCP provides a mock function with given fields: stream, reqAddr +func (_m *MockRequestHook) TCP(stream server.HyStream, reqAddr *string) ([]byte, error) { + ret := _m.Called(stream, reqAddr) + + if len(ret) == 0 { + panic("no return value specified for TCP") + } + + var r0 []byte + var r1 error + if rf, ok := ret.Get(0).(func(server.HyStream, *string) ([]byte, error)); ok { + return rf(stream, reqAddr) + } + if rf, ok := ret.Get(0).(func(server.HyStream, *string) []byte); ok { + r0 = rf(stream, reqAddr) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]byte) + } + } + + if rf, ok := ret.Get(1).(func(server.HyStream, *string) error); ok { + r1 = rf(stream, reqAddr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MockRequestHook_TCP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'TCP' +type MockRequestHook_TCP_Call struct { + *mock.Call +} + +// TCP is a helper method to define mock.On call +// - stream server.HyStream +// - reqAddr *string +func (_e *MockRequestHook_Expecter) TCP(stream interface{}, reqAddr interface{}) *MockRequestHook_TCP_Call { + return &MockRequestHook_TCP_Call{Call: _e.mock.On("TCP", stream, reqAddr)} +} + +func (_c *MockRequestHook_TCP_Call) Run(run func(stream server.HyStream, reqAddr *string)) *MockRequestHook_TCP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(server.HyStream), args[1].(*string)) + }) + return _c +} + +func (_c *MockRequestHook_TCP_Call) Return(_a0 []byte, _a1 error) *MockRequestHook_TCP_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *MockRequestHook_TCP_Call) RunAndReturn(run func(server.HyStream, *string) ([]byte, error)) *MockRequestHook_TCP_Call { + _c.Call.Return(run) + return _c +} + +// UDP provides a mock function with given fields: data, reqAddr +func (_m *MockRequestHook) UDP(data []byte, reqAddr *string) error { + ret := _m.Called(data, reqAddr) + + if len(ret) == 0 { + panic("no return value specified for UDP") + } + + var r0 error + if rf, ok := ret.Get(0).(func([]byte, *string) error); ok { + r0 = rf(data, reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockRequestHook_UDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UDP' +type MockRequestHook_UDP_Call struct { + *mock.Call +} + +// UDP is a helper method to define mock.On call +// - data []byte +// - reqAddr *string +func (_e *MockRequestHook_Expecter) UDP(data interface{}, reqAddr interface{}) *MockRequestHook_UDP_Call { + return &MockRequestHook_UDP_Call{Call: _e.mock.On("UDP", data, reqAddr)} +} + +func (_c *MockRequestHook_UDP_Call) Run(run func(data []byte, reqAddr *string)) *MockRequestHook_UDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte), args[1].(*string)) + }) + return _c +} + +func (_c *MockRequestHook_UDP_Call) Return(_a0 error) *MockRequestHook_UDP_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockRequestHook_UDP_Call) RunAndReturn(run func([]byte, *string) error) *MockRequestHook_UDP_Call { + _c.Call.Return(run) + return _c +} + +// NewMockRequestHook creates a new instance of MockRequestHook. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockRequestHook(t interface { + mock.TestingT + Cleanup(func()) +}) *MockRequestHook { + mock := &MockRequestHook{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_TrafficLogger.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_TrafficLogger.go new file mode 100644 index 0000000..92ed6ed --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_TrafficLogger.go @@ -0,0 +1,184 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import ( + server "github.com/apernet/hysteria/core/v2/server" + mock "github.com/stretchr/testify/mock" +) + +// MockTrafficLogger is an autogenerated mock type for the TrafficLogger type +type MockTrafficLogger struct { + mock.Mock +} + +type MockTrafficLogger_Expecter struct { + mock *mock.Mock +} + +func (_m *MockTrafficLogger) EXPECT() *MockTrafficLogger_Expecter { + return &MockTrafficLogger_Expecter{mock: &_m.Mock} +} + +// LogOnlineState provides a mock function with given fields: id, online +func (_m *MockTrafficLogger) LogOnlineState(id string, online bool) { + _m.Called(id, online) +} + +// MockTrafficLogger_LogOnlineState_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'LogOnlineState' +type MockTrafficLogger_LogOnlineState_Call struct { + *mock.Call +} + +// LogOnlineState is a helper method to define mock.On call +// - id string +// - online bool +func (_e *MockTrafficLogger_Expecter) LogOnlineState(id interface{}, online interface{}) *MockTrafficLogger_LogOnlineState_Call { + return &MockTrafficLogger_LogOnlineState_Call{Call: _e.mock.On("LogOnlineState", id, online)} +} + +func (_c *MockTrafficLogger_LogOnlineState_Call) Run(run func(id string, online bool)) *MockTrafficLogger_LogOnlineState_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string), args[1].(bool)) + }) + return _c +} + +func (_c *MockTrafficLogger_LogOnlineState_Call) Return() *MockTrafficLogger_LogOnlineState_Call { + _c.Call.Return() + return _c +} + +func (_c *MockTrafficLogger_LogOnlineState_Call) RunAndReturn(run func(string, bool)) *MockTrafficLogger_LogOnlineState_Call { + _c.Run(run) + return _c +} + +// LogTraffic provides a mock function with given fields: id, tx, rx +func (_m *MockTrafficLogger) LogTraffic(id string, tx uint64, rx uint64) bool { + ret := _m.Called(id, tx, rx) + + if len(ret) == 0 { + panic("no return value specified for LogTraffic") + } + + var r0 bool + if rf, ok := ret.Get(0).(func(string, uint64, uint64) bool); ok { + r0 = rf(id, tx, rx) + } else { + r0 = ret.Get(0).(bool) + } + + return r0 +} + +// MockTrafficLogger_LogTraffic_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'LogTraffic' +type MockTrafficLogger_LogTraffic_Call struct { + *mock.Call +} + +// LogTraffic is a helper method to define mock.On call +// - id string +// - tx uint64 +// - rx uint64 +func (_e *MockTrafficLogger_Expecter) LogTraffic(id interface{}, tx interface{}, rx interface{}) *MockTrafficLogger_LogTraffic_Call { + return &MockTrafficLogger_LogTraffic_Call{Call: _e.mock.On("LogTraffic", id, tx, rx)} +} + +func (_c *MockTrafficLogger_LogTraffic_Call) Run(run func(id string, tx uint64, rx uint64)) *MockTrafficLogger_LogTraffic_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string), args[1].(uint64), args[2].(uint64)) + }) + return _c +} + +func (_c *MockTrafficLogger_LogTraffic_Call) Return(ok bool) *MockTrafficLogger_LogTraffic_Call { + _c.Call.Return(ok) + return _c +} + +func (_c *MockTrafficLogger_LogTraffic_Call) RunAndReturn(run func(string, uint64, uint64) bool) *MockTrafficLogger_LogTraffic_Call { + _c.Call.Return(run) + return _c +} + +// TraceStream provides a mock function with given fields: stream, stats +func (_m *MockTrafficLogger) TraceStream(stream server.HyStream, stats *server.StreamStats) { + _m.Called(stream, stats) +} + +// MockTrafficLogger_TraceStream_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'TraceStream' +type MockTrafficLogger_TraceStream_Call struct { + *mock.Call +} + +// TraceStream is a helper method to define mock.On call +// - stream server.HyStream +// - stats *server.StreamStats +func (_e *MockTrafficLogger_Expecter) TraceStream(stream interface{}, stats interface{}) *MockTrafficLogger_TraceStream_Call { + return &MockTrafficLogger_TraceStream_Call{Call: _e.mock.On("TraceStream", stream, stats)} +} + +func (_c *MockTrafficLogger_TraceStream_Call) Run(run func(stream server.HyStream, stats *server.StreamStats)) *MockTrafficLogger_TraceStream_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(server.HyStream), args[1].(*server.StreamStats)) + }) + return _c +} + +func (_c *MockTrafficLogger_TraceStream_Call) Return() *MockTrafficLogger_TraceStream_Call { + _c.Call.Return() + return _c +} + +func (_c *MockTrafficLogger_TraceStream_Call) RunAndReturn(run func(server.HyStream, *server.StreamStats)) *MockTrafficLogger_TraceStream_Call { + _c.Run(run) + return _c +} + +// UntraceStream provides a mock function with given fields: stream +func (_m *MockTrafficLogger) UntraceStream(stream server.HyStream) { + _m.Called(stream) +} + +// MockTrafficLogger_UntraceStream_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UntraceStream' +type MockTrafficLogger_UntraceStream_Call struct { + *mock.Call +} + +// UntraceStream is a helper method to define mock.On call +// - stream server.HyStream +func (_e *MockTrafficLogger_Expecter) UntraceStream(stream interface{}) *MockTrafficLogger_UntraceStream_Call { + return &MockTrafficLogger_UntraceStream_Call{Call: _e.mock.On("UntraceStream", stream)} +} + +func (_c *MockTrafficLogger_UntraceStream_Call) Run(run func(stream server.HyStream)) *MockTrafficLogger_UntraceStream_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(server.HyStream)) + }) + return _c +} + +func (_c *MockTrafficLogger_UntraceStream_Call) Return() *MockTrafficLogger_UntraceStream_Call { + _c.Call.Return() + return _c +} + +func (_c *MockTrafficLogger_UntraceStream_Call) RunAndReturn(run func(server.HyStream)) *MockTrafficLogger_UntraceStream_Call { + _c.Run(run) + return _c +} + +// NewMockTrafficLogger creates a new instance of MockTrafficLogger. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockTrafficLogger(t interface { + mock.TestingT + Cleanup(func()) +}) *MockTrafficLogger { + mock := &MockTrafficLogger{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/mocks/mock_UDPConn.go b/third_party/hysteria-core/internal/integration_tests/mocks/mock_UDPConn.go new file mode 100644 index 0000000..965edc0 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/mocks/mock_UDPConn.go @@ -0,0 +1,197 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package mocks + +import mock "github.com/stretchr/testify/mock" + +// MockUDPConn is an autogenerated mock type for the UDPConn type +type MockUDPConn struct { + mock.Mock +} + +type MockUDPConn_Expecter struct { + mock *mock.Mock +} + +func (_m *MockUDPConn) EXPECT() *MockUDPConn_Expecter { + return &MockUDPConn_Expecter{mock: &_m.Mock} +} + +// Close provides a mock function with no fields +func (_m *MockUDPConn) Close() error { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for Close") + } + + var r0 error + if rf, ok := ret.Get(0).(func() error); ok { + r0 = rf() + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockUDPConn_Close_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Close' +type MockUDPConn_Close_Call struct { + *mock.Call +} + +// Close is a helper method to define mock.On call +func (_e *MockUDPConn_Expecter) Close() *MockUDPConn_Close_Call { + return &MockUDPConn_Close_Call{Call: _e.mock.On("Close")} +} + +func (_c *MockUDPConn_Close_Call) Run(run func()) *MockUDPConn_Close_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *MockUDPConn_Close_Call) Return(_a0 error) *MockUDPConn_Close_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockUDPConn_Close_Call) RunAndReturn(run func() error) *MockUDPConn_Close_Call { + _c.Call.Return(run) + return _c +} + +// ReadFrom provides a mock function with given fields: b +func (_m *MockUDPConn) ReadFrom(b []byte) (int, string, error) { + ret := _m.Called(b) + + if len(ret) == 0 { + panic("no return value specified for ReadFrom") + } + + var r0 int + var r1 string + var r2 error + if rf, ok := ret.Get(0).(func([]byte) (int, string, error)); ok { + return rf(b) + } + if rf, ok := ret.Get(0).(func([]byte) int); ok { + r0 = rf(b) + } else { + r0 = ret.Get(0).(int) + } + + if rf, ok := ret.Get(1).(func([]byte) string); ok { + r1 = rf(b) + } else { + r1 = ret.Get(1).(string) + } + + if rf, ok := ret.Get(2).(func([]byte) error); ok { + r2 = rf(b) + } else { + r2 = ret.Error(2) + } + + return r0, r1, r2 +} + +// MockUDPConn_ReadFrom_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ReadFrom' +type MockUDPConn_ReadFrom_Call struct { + *mock.Call +} + +// ReadFrom is a helper method to define mock.On call +// - b []byte +func (_e *MockUDPConn_Expecter) ReadFrom(b interface{}) *MockUDPConn_ReadFrom_Call { + return &MockUDPConn_ReadFrom_Call{Call: _e.mock.On("ReadFrom", b)} +} + +func (_c *MockUDPConn_ReadFrom_Call) Run(run func(b []byte)) *MockUDPConn_ReadFrom_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte)) + }) + return _c +} + +func (_c *MockUDPConn_ReadFrom_Call) Return(_a0 int, _a1 string, _a2 error) *MockUDPConn_ReadFrom_Call { + _c.Call.Return(_a0, _a1, _a2) + return _c +} + +func (_c *MockUDPConn_ReadFrom_Call) RunAndReturn(run func([]byte) (int, string, error)) *MockUDPConn_ReadFrom_Call { + _c.Call.Return(run) + return _c +} + +// WriteTo provides a mock function with given fields: b, addr +func (_m *MockUDPConn) WriteTo(b []byte, addr string) (int, error) { + ret := _m.Called(b, addr) + + if len(ret) == 0 { + panic("no return value specified for WriteTo") + } + + var r0 int + var r1 error + if rf, ok := ret.Get(0).(func([]byte, string) (int, error)); ok { + return rf(b, addr) + } + if rf, ok := ret.Get(0).(func([]byte, string) int); ok { + r0 = rf(b, addr) + } else { + r0 = ret.Get(0).(int) + } + + if rf, ok := ret.Get(1).(func([]byte, string) error); ok { + r1 = rf(b, addr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MockUDPConn_WriteTo_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'WriteTo' +type MockUDPConn_WriteTo_Call struct { + *mock.Call +} + +// WriteTo is a helper method to define mock.On call +// - b []byte +// - addr string +func (_e *MockUDPConn_Expecter) WriteTo(b interface{}, addr interface{}) *MockUDPConn_WriteTo_Call { + return &MockUDPConn_WriteTo_Call{Call: _e.mock.On("WriteTo", b, addr)} +} + +func (_c *MockUDPConn_WriteTo_Call) Run(run func(b []byte, addr string)) *MockUDPConn_WriteTo_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte), args[1].(string)) + }) + return _c +} + +func (_c *MockUDPConn_WriteTo_Call) Return(_a0 int, _a1 error) *MockUDPConn_WriteTo_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *MockUDPConn_WriteTo_Call) RunAndReturn(run func([]byte, string) (int, error)) *MockUDPConn_WriteTo_Call { + _c.Call.Return(run) + return _c +} + +// NewMockUDPConn creates a new instance of MockUDPConn. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockUDPConn(t interface { + mock.TestingT + Cleanup(func()) +}) *MockUDPConn { + mock := &MockUDPConn{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/internal/integration_tests/smoke_test.go b/third_party/hysteria-core/internal/integration_tests/smoke_test.go new file mode 100644 index 0000000..ef14991 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/smoke_test.go @@ -0,0 +1,283 @@ +package integration_tests + +import ( + "io" + "net" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + + "github.com/apernet/hysteria/core/v2/client" + coreErrs "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" +) + +// Smoke tests that act as a sanity check for client & server to ensure they can talk to each other correctly. + +// TestClientNoServer tests how the client handles a server address it cannot connect to. +// NewClient should return a ConnectError. +func TestClientNoServer(t *testing.T) { + c, _, err := client.NewClient(&client.Config{ + ServerAddr: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 55666}, + }) + assert.Nil(t, c) + _, ok := err.(coreErrs.ConnectError) + assert.True(t, ok) +} + +// TestClientServerBadAuth tests two things: +// - The server uses Authenticator when a client connects. +// - How the client handles failed authentication. +func TestClientServerBadAuth(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, "badpassword", uint64(0)).Return(false, "").Once() + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + Auth: "badpassword", + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.Nil(t, c) + _, ok := err.(coreErrs.AuthError) + assert.True(t, ok) +} + +// TestClientServerUDPDisabled tests how the client handles a server that does not support UDP. +// UDP should return a DialError. +func TestClientServerUDPDisabled(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + DisableUDP: true, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + conn, err := c.UDP() + assert.Nil(t, conn) + _, ok := err.(coreErrs.DialError) + assert.True(t, ok) +} + +// TestClientServerTCPEcho tests TCP forwarding using a TCP echo server. +func TestClientServerTCPEcho(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create TCP echo server + echoAddr := "127.0.0.1:22333" + echoListener, err := net.Listen("tcp", echoAddr) + assert.NoError(t, err) + echoServer := &tcpEchoServer{Listener: echoListener} + defer echoServer.Close() + go echoServer.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + // Dial TCP + conn, err := c.TCP(echoAddr) + assert.NoError(t, err) + defer conn.Close() + + // Send and receive data + sData := []byte("hello world") + _, err = conn.Write(sData) + assert.NoError(t, err) + rData := make([]byte, len(sData)) + _, err = io.ReadFull(conn, rData) + assert.NoError(t, err) + assert.Equal(t, sData, rData) +} + +// TestClientServerUDPEcho tests UDP forwarding using a UDP echo server. +func TestClientServerUDPEcho(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create UDP echo server + echoAddr := "127.0.0.1:22333" + echoConn, err := net.ListenPacket("udp", echoAddr) + assert.NoError(t, err) + echoServer := &udpEchoServer{Conn: echoConn} + defer echoServer.Close() + go echoServer.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + // Listen UDP + conn, err := c.UDP() + assert.NoError(t, err) + defer conn.Close() + + // Send and receive data + sData := []byte("hello world") + err = conn.Send(sData, echoAddr) + assert.NoError(t, err) + rData, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, sData, rData) + assert.Equal(t, echoAddr, rAddr) +} + +// TestClientServerHandshakeInfo tests that the client returns the correct handshake info. +func TestClientServerHandshakeInfo(t *testing.T) { + // Create server 1, UDP enabled, unlimited bandwidth + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + }) + assert.NoError(t, err) + go s.Serve() + + // Create client 1, with specified tx bandwidth + c, info, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + BandwidthConfig: client.BandwidthConfig{ + MaxTx: 123456, + }, + }) + assert.NoError(t, err) + assert.Equal(t, &client.HandshakeInfo{ + UDPEnabled: true, + Tx: 123456, + ServerAddr: udpAddr, + }, info) + + // Close server 1 and client 1 + _ = s.Close() + _ = c.Close() + + // Create server 2, UDP disabled, limited rx bandwidth + udpConn, udpAddr, err = serverConn() + assert.NoError(t, err) + s, err = server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + BandwidthConfig: server.BandwidthConfig{ + MaxRx: 100000, + }, + DisableUDP: true, + Authenticator: auth, + }) + assert.NoError(t, err) + go s.Serve() + + // Create client 2, with specified tx bandwidth + c, info, err = client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + BandwidthConfig: client.BandwidthConfig{ + MaxTx: 123456, + }, + }) + assert.NoError(t, err) + assert.Equal(t, &client.HandshakeInfo{ + UDPEnabled: false, + Tx: 100000, + ServerAddr: udpAddr, + }, info) + + // Close server 2 and client 2 + _ = s.Close() + _ = c.Close() + + // Create server 3, UDP enabled, ignore client bandwidth + udpConn, udpAddr, err = serverConn() + assert.NoError(t, err) + s, err = server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + IgnoreClientBandwidth: true, + Authenticator: auth, + }) + assert.NoError(t, err) + go s.Serve() + + // Create client 3, with specified tx bandwidth + c, info, err = client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + BandwidthConfig: client.BandwidthConfig{ + MaxTx: 123456, + }, + }) + assert.NoError(t, err) + assert.Equal(t, &client.HandshakeInfo{ + UDPEnabled: true, + Tx: 0, + ServerAddr: udpAddr, + }, info) + + // Close server 3 and client 3 + _ = s.Close() + _ = c.Close() +} diff --git a/third_party/hysteria-core/internal/integration_tests/stress_test.go b/third_party/hysteria-core/internal/integration_tests/stress_test.go new file mode 100644 index 0000000..fc98847 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/stress_test.go @@ -0,0 +1,263 @@ +package integration_tests + +import ( + "context" + "crypto/rand" + "fmt" + "io" + "net" + "os" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "golang.org/x/time/rate" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" +) + +type tcpStressor struct { + DialFunc func() (net.Conn, error) + Size int + Parallel int + Iterations int +} + +func (s *tcpStressor) Run(t *testing.T) { + // Make some random data + sData := make([]byte, s.Size) + _, err := rand.Read(sData) + assert.NoError(t, err) + + // Run iterations + for i := 0; i < s.Iterations; i++ { + var wg sync.WaitGroup + errChan := make(chan error, s.Parallel) + for j := 0; j < s.Parallel; j++ { + wg.Add(1) + go func() { + defer wg.Done() + + conn, err := s.DialFunc() + if err != nil { + errChan <- err + return + } + defer conn.Close() + go conn.Write(sData) + + rData := make([]byte, len(sData)) + _, err = io.ReadFull(conn, rData) + if err != nil { + errChan <- err + return + } + }() + } + wg.Wait() + + assert.Empty(t, errChan) + } +} + +type udpStressor struct { + ListenFunc func() (client.HyUDPConn, error) + ServerAddr string + Size int + Count int + Parallel int + Iterations int +} + +func (s *udpStressor) Run(t *testing.T) { + // Make some random data + sData := make([]byte, s.Size) + _, err := rand.Read(sData) + assert.NoError(t, err) + + // Due to UDP's unreliability, we need to limit the rate of sending + // to reduce packet loss. This is hardcoded to 1 MiB/s for now. + limiter := rate.NewLimiter(1048576, 1048576) + + // Run iterations + for i := 0; i < s.Iterations; i++ { + var wg sync.WaitGroup + errChan := make(chan error, s.Parallel) + for j := 0; j < s.Parallel; j++ { + wg.Add(1) + go func() { + defer wg.Done() + + conn, err := s.ListenFunc() + if err != nil { + errChan <- err + return + } + defer conn.Close() + go func() { + // Sending routine + for i := 0; i < s.Count; i++ { + _ = limiter.WaitN(context.Background(), len(sData)) + _ = conn.Send(sData, s.ServerAddr) + } + }() + + minCount := s.Count * 8 / 10 // Tolerate 20% packet loss + for i := 0; i < minCount; i++ { + rData, _, err := conn.Receive() + if err != nil { + errChan <- err + return + } + if len(rData) != len(sData) { + errChan <- fmt.Errorf("incomplete data received: %d/%d bytes", len(rData), len(sData)) + return + } + } + }() + } + wg.Wait() + + assert.Empty(t, errChan) + } +} + +func TestClientServerTCPStress(t *testing.T) { + if os.Getenv("AUTOCAR_RUN_UPSTREAM_STRESS") != "1" { + t.Skip("set AUTOCAR_RUN_UPSTREAM_STRESS=1 to run multi-gigabyte upstream stress cases") + } + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + MaxTCPHandlers: 1024, + MaxClientTCPHandlers: 1024, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create TCP echo server + echoAddr := "127.0.0.1:22333" + echoListener, err := net.Listen("tcp", echoAddr) + assert.NoError(t, err) + echoServer := &tcpEchoServer{Listener: echoListener} + defer echoServer.Close() + go echoServer.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + dialFunc := func() (net.Conn, error) { + return c.TCP(echoAddr) + } + + t.Run("Single 500m", (&tcpStressor{DialFunc: dialFunc, Size: 524288000, Parallel: 1, Iterations: 1}).Run) + + t.Run("Sequential 1000x1m", (&tcpStressor{DialFunc: dialFunc, Size: 1048576, Parallel: 1, Iterations: 1000}).Run) + t.Run("Sequential 10000x100k", (&tcpStressor{DialFunc: dialFunc, Size: 102400, Parallel: 1, Iterations: 10000}).Run) + + t.Run("Parallel 100x10m", (&tcpStressor{DialFunc: dialFunc, Size: 10485760, Parallel: 100, Iterations: 1}).Run) + t.Run("Parallel 1000x1m", (&tcpStressor{DialFunc: dialFunc, Size: 1048576, Parallel: 1000, Iterations: 1}).Run) +} + +func TestClientServerUDPStress(t *testing.T) { + if os.Getenv("AUTOCAR_RUN_UPSTREAM_STRESS") != "1" { + t.Skip("set AUTOCAR_RUN_UPSTREAM_STRESS=1 to run long lossy-UDP upstream stress cases") + } + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + MaxUDPSessions: 1024, + MaxClientUDPSessions: 1024, + Authenticator: auth, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create UDP echo server + echoAddr := "127.0.0.1:22333" + echoConn, err := net.ListenPacket("udp", echoAddr) + assert.NoError(t, err) + echoServer := &udpEchoServer{Conn: echoConn} + defer echoServer.Close() + go echoServer.Serve() + + // Create client + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + t.Run("Single 1000x100b", (&udpStressor{ + ListenFunc: c.UDP, + ServerAddr: echoAddr, + Size: 100, + Count: 1000, + Parallel: 1, + Iterations: 1, + }).Run) + t.Run("Single 1000x3k", (&udpStressor{ + ListenFunc: c.UDP, + ServerAddr: echoAddr, + Size: 3000, + Count: 1000, + Parallel: 1, + Iterations: 1, + }).Run) + + t.Run("5 Sequential 1000x100b", (&udpStressor{ + ListenFunc: c.UDP, + ServerAddr: echoAddr, + Size: 100, + Count: 1000, + Parallel: 1, + Iterations: 5, + }).Run) + t.Run("5 Sequential 200x3k", (&udpStressor{ + ListenFunc: c.UDP, + ServerAddr: echoAddr, + Size: 3000, + Count: 200, + Parallel: 1, + Iterations: 5, + }).Run) + + t.Run("2 Sequential 5 Parallel 1000x100b", (&udpStressor{ + ListenFunc: c.UDP, + ServerAddr: echoAddr, + Size: 100, + Count: 1000, + Parallel: 5, + Iterations: 2, + }).Run) + t.Run("2 Sequential 5 Parallel 200x3k", (&udpStressor{ + ListenFunc: c.UDP, + ServerAddr: echoAddr, + Size: 3000, + Count: 200, + Parallel: 5, + Iterations: 2, + }).Run) +} diff --git a/third_party/hysteria-core/internal/integration_tests/trafficlogger_test.go b/third_party/hysteria-core/internal/integration_tests/trafficlogger_test.go new file mode 100644 index 0000000..841f4ff --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/trafficlogger_test.go @@ -0,0 +1,180 @@ +package integration_tests + +import ( + "io" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" +) + +// TestClientServerTrafficLoggerTCP tests that the traffic logger is correctly called for TCP connections, +// and that the client is disconnected when the traffic logger returns false. +func TestClientServerTrafficLoggerTCP(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + serverOb := mocks.NewMockOutbound(t) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + trafficLogger := mocks.NewMockTrafficLogger(t) + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Outbound: serverOb, + Authenticator: auth, + TrafficLogger: trafficLogger, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + trafficLogger.EXPECT().LogOnlineState("nobody", true).Return().Once() + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + addr := "dontcare.cc:4455" + + sobConn := mocks.NewMockConn(t) + sobConnCh := make(chan []byte, 1) + sobConnChCloseFunc := sync.OnceFunc(func() { close(sobConnCh) }) + sobConn.EXPECT().Read(mock.Anything).RunAndReturn(func(bs []byte) (int, error) { + b := <-sobConnCh + if b == nil { + return 0, io.EOF + } else { + return copy(bs, b), nil + } + }) + sobConn.EXPECT().Close().RunAndReturn(func() error { + sobConnChCloseFunc() + return nil + }) + serverOb.EXPECT().TCP(addr).Return(sobConn, nil).Once() + trafficLogger.EXPECT().TraceStream(mock.Anything, mock.Anything).Return().Once() + + conn, err := c.TCP(addr) + assert.NoError(t, err) + + // Client reads from server + trafficLogger.EXPECT().LogTraffic("nobody", uint64(0), uint64(11)).Return(true).Once() + sobConnCh <- []byte("knock knock") + buf := make([]byte, 100) + n, err := conn.Read(buf) + assert.NoError(t, err) + assert.Equal(t, 11, n) + assert.Equal(t, "knock knock", string(buf[:n])) + + // Client writes to server + trafficLogger.EXPECT().LogTraffic("nobody", uint64(12), uint64(0)).Return(true).Once() + sobConn.EXPECT().Write([]byte("who is there")).Return(12, nil).Once() + n, err = conn.Write([]byte("who is there")) + assert.NoError(t, err) + assert.Equal(t, 12, n) + time.Sleep(1 * time.Second) // Need some time for the server to receive the data + + // Client reads from server again but blocked + trafficLogger.EXPECT().UntraceStream(mock.Anything).Return().Once() + trafficLogger.EXPECT().LogTraffic("nobody", uint64(0), uint64(4)).Return(false).Once() + trafficLogger.EXPECT().LogOnlineState("nobody", false).Return().Once() + sobConnCh <- []byte("nope") + n, err = conn.Read(buf) + assert.Zero(t, n) + assert.Error(t, err) + + // The client should be disconnected + _, err = c.TCP("whatever") + assert.Error(t, err) +} + +// TestClientServerTrafficLoggerUDP tests that the traffic logger is correctly called for UDP sessions, +// and that the client is disconnected when the traffic logger returns false. +func TestClientServerTrafficLoggerUDP(t *testing.T) { + // Create server + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + serverOb := mocks.NewMockOutbound(t) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + trafficLogger := mocks.NewMockTrafficLogger(t) + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Outbound: serverOb, + Authenticator: auth, + TrafficLogger: trafficLogger, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + // Create client + trafficLogger.EXPECT().LogOnlineState("nobody", true).Return().Once() + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + addr := "shady.org:43211" + + sobConn := mocks.NewMockUDPConn(t) + sobConnCh := make(chan []byte, 1) + sobConnChCloseFunc := sync.OnceFunc(func() { close(sobConnCh) }) + sobConn.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(bs []byte) (int, string, error) { + b := <-sobConnCh + if b == nil { + return 0, "", io.EOF + } else { + return copy(bs, b), addr, nil + } + }) + sobConn.EXPECT().Close().RunAndReturn(func() error { + sobConnChCloseFunc() + return nil + }) + serverOb.EXPECT().UDP(addr).Return(sobConn, nil).Once() + + conn, err := c.UDP() + assert.NoError(t, err) + + // Client writes to server + trafficLogger.EXPECT().LogTraffic("nobody", uint64(9), uint64(0)).Return(true).Once() + sobConn.EXPECT().WriteTo([]byte("small sad"), addr).Return(9, nil).Once() + err = conn.Send([]byte("small sad"), addr) + assert.NoError(t, err) + time.Sleep(1 * time.Second) // Need some time for the server to receive the data + + // Client reads from server + trafficLogger.EXPECT().LogTraffic("nobody", uint64(0), uint64(7)).Return(true).Once() + sobConnCh <- []byte("big mad") + bs, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, rAddr, addr) + assert.Equal(t, "big mad", string(bs)) + + // Client reads from server again but blocked + trafficLogger.EXPECT().LogTraffic("nobody", uint64(0), uint64(4)).Return(false).Once() + trafficLogger.EXPECT().LogOnlineState("nobody", false).Return().Once() + sobConnCh <- []byte("nope") + bs, rAddr, err = conn.Receive() + assert.Equal(t, err, io.EOF) + assert.Empty(t, rAddr) + assert.Empty(t, bs) + + // The client should be disconnected + _, err = c.UDP() + assert.Error(t, err) +} diff --git a/third_party/hysteria-core/internal/integration_tests/udp_acl_test.go b/third_party/hysteria-core/internal/integration_tests/udp_acl_test.go new file mode 100644 index 0000000..b8bda01 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/udp_acl_test.go @@ -0,0 +1,177 @@ +package integration_tests + +import ( + "errors" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +type gatedOutbound struct { + blocked string + checkCalls atomic.Int32 + dialedAddrs atomic.Int32 +} + +func (o *gatedOutbound) TCP(reqAddr string) (net.Conn, error) { + return net.Dial("tcp", reqAddr) +} + +func (o *gatedOutbound) UDP(reqAddr string) (server.UDPConn, error) { + if reqAddr == o.blocked { + return nil, errors.New("rejected") + } + o.dialedAddrs.Add(1) + c, err := net.ListenUDP("udp", nil) + if err != nil { + return nil, err + } + return &gatedUDPConn{UDPConn: c}, nil +} + +func (o *gatedOutbound) CheckUDP(reqAddr string) error { + o.checkCalls.Add(1) + if reqAddr == o.blocked { + return errors.New("rejected") + } + return nil +} + +type gatedUDPConn struct { + *net.UDPConn +} + +func (c *gatedUDPConn) ReadFrom(b []byte) (int, string, error) { + n, addr, err := c.UDPConn.ReadFrom(b) + if addr != nil { + return n, addr.String(), err + } + return n, "", err +} + +func (c *gatedUDPConn) WriteTo(b []byte, addr string) (int, error) { + uAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return 0, err + } + return c.UDPConn.WriteTo(b, uAddr) +} + +func TestClientServerUDPACLBypass(t *testing.T) { + const allowed, blocked = "127.0.0.1:22444", "127.0.0.1:22445" + ob := &gatedOutbound{blocked: blocked} + + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + Outbound: ob, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + allowedConn, err := net.ListenPacket("udp", allowed) + assert.NoError(t, err) + defer allowedConn.Close() + go (&udpEchoServer{Conn: allowedConn}).Serve() + + blockedConn, err := net.ListenPacket("udp", blocked) + assert.NoError(t, err) + defer blockedConn.Close() + go (&udpEchoServer{Conn: blockedConn}).Serve() + + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + conn, err := c.UDP() + assert.NoError(t, err) + defer conn.Close() + + assert.NoError(t, conn.Send([]byte("hello"), allowed)) + rData, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, []byte("hello"), rData) + assert.Equal(t, allowed, rAddr) + + assert.NoError(t, conn.Send([]byte("ssrf"), blocked)) + + done := make(chan struct{}) + var leakedAddr string + go func() { + _, addr, err := conn.Receive() + if err == nil { + leakedAddr = addr + } + close(done) + }() + select { + case <-done: + assert.NotEqual(t, blocked, leakedAddr, "ACL bypass: blocked destination relayed") + case <-time.After(500 * time.Millisecond): + } + + assert.GreaterOrEqual(t, ob.checkCalls.Load(), int32(1), "CheckUDP not invoked for subsequent packet") + assert.Equal(t, int32(1), ob.dialedAddrs.Load(), "outbound dial must happen only on first allowed destination") +} + +func TestClientServerUDPACLMultiDestAllowed(t *testing.T) { + const dest1, dest2 = "127.0.0.1:22448", "127.0.0.1:22449" + ob := &gatedOutbound{blocked: ""} + + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + Outbound: ob, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + for _, addr := range []string{dest1, dest2} { + ec, err := net.ListenPacket("udp", addr) + assert.NoError(t, err) + defer ec.Close() + go (&udpEchoServer{Conn: ec}).Serve() + } + + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + conn, err := c.UDP() + assert.NoError(t, err) + defer conn.Close() + + for _, addr := range []string{dest1, dest2} { + assert.NoError(t, conn.Send([]byte("hi"), addr)) + rData, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, []byte("hi"), rData) + assert.Equal(t, addr, rAddr) + } +} diff --git a/third_party/hysteria-core/internal/integration_tests/utils_test.go b/third_party/hysteria-core/internal/integration_tests/utils_test.go new file mode 100644 index 0000000..bea8985 --- /dev/null +++ b/third_party/hysteria-core/internal/integration_tests/utils_test.go @@ -0,0 +1,104 @@ +package integration_tests + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "io" + "math/big" + "net" + "sync" + "time" + + "github.com/apernet/hysteria/core/v2/server" +) + +// This file provides utilities for the integration tests. + +var testCertificate = sync.OnceValue(func() tls.Certificate { + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + panic(err) + } + now := time.Now() + template := &x509.Certificate{ + SerialNumber: big.NewInt(now.UnixNano()), + Subject: pkix.Name{CommonName: "localhost"}, + NotBefore: now.Add(-time.Minute), + NotAfter: now.Add(24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + BasicConstraintsValid: true, + DNSNames: []string{"localhost"}, + } + der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + if err != nil { + panic(err) + } + return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key} +}) + +func serverTLSConfig() server.TLSConfig { + return server.TLSConfig{ + Certificates: []tls.Certificate{testCertificate()}, + } +} + +func serverConn() (net.PacketConn, net.Addr, error) { + udpAddr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 14514} + udpConn, err := net.ListenUDP("udp", udpAddr) + if err != nil { + return nil, nil, err + } + return udpConn, udpAddr, nil +} + +// tcpEchoServer is a TCP server that echoes what it reads from the connection. +// It will never actively close the connection. +type tcpEchoServer struct { + Listener net.Listener +} + +func (s *tcpEchoServer) Serve() error { + for { + conn, err := s.Listener.Accept() + if err != nil { + return err + } + go func() { + _, _ = io.Copy(conn, conn) + _ = conn.Close() + }() + } +} + +func (s *tcpEchoServer) Close() error { + return s.Listener.Close() +} + +// udpEchoServer is a UDP server that echoes what it reads from the connection. +// It will never actively close the connection. +type udpEchoServer struct { + Conn net.PacketConn +} + +func (s *udpEchoServer) Serve() error { + buf := make([]byte, 65536) + for { + n, addr, err := s.Conn.ReadFrom(buf) + if err != nil { + return err + } + _, err = s.Conn.WriteTo(buf[:n], addr) + if err != nil { + return err + } + } +} + +func (s *udpEchoServer) Close() error { + return s.Conn.Close() +} diff --git a/third_party/hysteria-core/internal/pmtud/avail.go b/third_party/hysteria-core/internal/pmtud/avail.go new file mode 100644 index 0000000..cd7afd0 --- /dev/null +++ b/third_party/hysteria-core/internal/pmtud/avail.go @@ -0,0 +1,7 @@ +//go:build linux || windows || darwin + +package pmtud + +const ( + DisablePathMTUDiscovery = false +) diff --git a/third_party/hysteria-core/internal/pmtud/unavail.go b/third_party/hysteria-core/internal/pmtud/unavail.go new file mode 100644 index 0000000..917b973 --- /dev/null +++ b/third_party/hysteria-core/internal/pmtud/unavail.go @@ -0,0 +1,13 @@ +//go:build !linux && !windows && !darwin + +package pmtud + +// quic-go's MTU detection is enabled by default on all platforms. +// However, it only actually sets the DF bit on 3 supported platforms (Windows, macOS, Linux). +// As a result, on other platforms, probe packets that should never be fragmented will still +// be fragmented and transmitted. So we're only enabling it for platforms where we've verified +// its functionality for now. + +const ( + DisablePathMTUDiscovery = true +) diff --git a/third_party/hysteria-core/internal/protocol/http.go b/third_party/hysteria-core/internal/protocol/http.go new file mode 100644 index 0000000..abcc1a4 --- /dev/null +++ b/third_party/hysteria-core/internal/protocol/http.go @@ -0,0 +1,68 @@ +package protocol + +import ( + "net/http" + "strconv" +) + +const ( + URLHost = "hysteria" + URLPath = "/auth" + + RequestHeaderAuth = "Hysteria-Auth" + ResponseHeaderUDPEnabled = "Hysteria-UDP" + CommonHeaderCCRX = "Hysteria-CC-RX" + CommonHeaderPadding = "Hysteria-Padding" + + StatusAuthOK = 233 +) + +// AuthRequest is what client sends to server for authentication. +type AuthRequest struct { + Auth string + Rx uint64 // 0 = unknown, client asks server to use bandwidth detection +} + +// AuthResponse is what server sends to client when authentication is passed. +type AuthResponse struct { + UDPEnabled bool + Rx uint64 // 0 = unlimited + RxAuto bool // true = server asks client to use bandwidth detection +} + +func AuthRequestFromHeader(h http.Header) AuthRequest { + rx, _ := strconv.ParseUint(h.Get(CommonHeaderCCRX), 10, 64) + return AuthRequest{ + Auth: h.Get(RequestHeaderAuth), + Rx: rx, + } +} + +func AuthRequestToHeader(h http.Header, req AuthRequest) { + h.Set(RequestHeaderAuth, req.Auth) + h.Set(CommonHeaderCCRX, strconv.FormatUint(req.Rx, 10)) + h.Set(CommonHeaderPadding, authRequestPadding.String()) +} + +func AuthResponseFromHeader(h http.Header) AuthResponse { + resp := AuthResponse{} + resp.UDPEnabled, _ = strconv.ParseBool(h.Get(ResponseHeaderUDPEnabled)) + rxStr := h.Get(CommonHeaderCCRX) + if rxStr == "auto" { + // Special case for server requesting client to use bandwidth detection + resp.RxAuto = true + } else { + resp.Rx, _ = strconv.ParseUint(rxStr, 10, 64) + } + return resp +} + +func AuthResponseToHeader(h http.Header, resp AuthResponse) { + h.Set(ResponseHeaderUDPEnabled, strconv.FormatBool(resp.UDPEnabled)) + if resp.RxAuto { + h.Set(CommonHeaderCCRX, "auto") + } else { + h.Set(CommonHeaderCCRX, strconv.FormatUint(resp.Rx, 10)) + } + h.Set(CommonHeaderPadding, authResponsePadding.String()) +} diff --git a/third_party/hysteria-core/internal/protocol/padding.go b/third_party/hysteria-core/internal/protocol/padding.go new file mode 100644 index 0000000..9895cdc --- /dev/null +++ b/third_party/hysteria-core/internal/protocol/padding.go @@ -0,0 +1,31 @@ +package protocol + +import ( + "math/rand" +) + +const ( + paddingChars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" +) + +// padding specifies a half-open range [Min, Max). +type padding struct { + Min int + Max int +} + +func (p padding) String() string { + n := p.Min + rand.Intn(p.Max-p.Min) + bs := make([]byte, n) + for i := range bs { + bs[i] = paddingChars[rand.Intn(len(paddingChars))] + } + return string(bs) +} + +var ( + authRequestPadding = padding{Min: 256, Max: 2048} + authResponsePadding = padding{Min: 256, Max: 2048} + tcpRequestPadding = padding{Min: 64, Max: 512} + tcpResponsePadding = padding{Min: 128, Max: 1024} +) diff --git a/third_party/hysteria-core/internal/protocol/proxy.go b/third_party/hysteria-core/internal/protocol/proxy.go new file mode 100644 index 0000000..8448511 --- /dev/null +++ b/third_party/hysteria-core/internal/protocol/proxy.go @@ -0,0 +1,261 @@ +package protocol + +import ( + "bytes" + "encoding/binary" + "fmt" + "io" + + "github.com/apernet/hysteria/core/v2/errors" + + "github.com/apernet/quic-go/quicvarint" +) + +const ( + FrameTypeTCPRequest = 0x401 + + // Max length values are for preventing DoS attacks + + MaxAddressLength = 2048 + MaxMessageLength = 2048 + MaxPaddingLength = 4096 + + MaxDatagramFrameSize = 1200 + MaxUDPSize = 4096 + // MaxUDPMessageSize includes the largest accepted payload, address and + // fixed/varint framing. Send buffers must use this size so a maximum UDP + // payload reaches QUIC, which can then return its real datagram limit and + // trigger fragmentation instead of being dropped during serialization. + MaxUDPMessageSize = 8 + 2 + MaxAddressLength + MaxUDPSize + + maxVarInt1 = 63 + maxVarInt2 = 16383 + maxVarInt4 = 1073741823 + maxVarInt8 = 4611686018427387903 +) + +// TCPRequest format: +// 0x401 (QUIC varint) +// Address length (QUIC varint) +// Address (bytes) +// Padding length (QUIC varint) +// Padding (bytes) + +func ReadTCPRequest(r io.Reader) (string, error) { + bReader := quicvarint.NewReader(r) + addrLen, err := quicvarint.Read(bReader) + if err != nil { + return "", err + } + if addrLen == 0 || addrLen > MaxAddressLength { + return "", errors.ProtocolError{Message: "invalid address length"} + } + addrBuf := make([]byte, addrLen) + _, err = io.ReadFull(r, addrBuf) + if err != nil { + return "", err + } + paddingLen, err := quicvarint.Read(bReader) + if err != nil { + return "", err + } + if paddingLen > MaxPaddingLength { + return "", errors.ProtocolError{Message: "invalid padding length"} + } + if paddingLen > 0 { + _, err = io.CopyN(io.Discard, r, int64(paddingLen)) + if err != nil { + return "", err + } + } + return string(addrBuf), nil +} + +func WriteTCPRequest(w io.Writer, addr string) error { + padding := tcpRequestPadding.String() + paddingLen := len(padding) + addrLen := len(addr) + sz := int(quicvarint.Len(FrameTypeTCPRequest)) + + int(quicvarint.Len(uint64(addrLen))) + addrLen + + int(quicvarint.Len(uint64(paddingLen))) + paddingLen + buf := make([]byte, sz) + i := varintPut(buf, FrameTypeTCPRequest) + i += varintPut(buf[i:], uint64(addrLen)) + i += copy(buf[i:], addr) + i += varintPut(buf[i:], uint64(paddingLen)) + copy(buf[i:], padding) + _, err := w.Write(buf) + return err +} + +// TCPResponse format: +// Status (byte, 0=ok, 1=error) +// Message length (QUIC varint) +// Message (bytes) +// Padding length (QUIC varint) +// Padding (bytes) + +func ReadTCPResponse(r io.Reader) (bool, string, error) { + var status [1]byte + if _, err := io.ReadFull(r, status[:]); err != nil { + return false, "", err + } + bReader := quicvarint.NewReader(r) + msgLen, err := quicvarint.Read(bReader) + if err != nil { + return false, "", err + } + if msgLen > MaxMessageLength { + return false, "", errors.ProtocolError{Message: "invalid message length"} + } + var msgBuf []byte + // No message is fine + if msgLen > 0 { + msgBuf = make([]byte, msgLen) + _, err = io.ReadFull(r, msgBuf) + if err != nil { + return false, "", err + } + } + paddingLen, err := quicvarint.Read(bReader) + if err != nil { + return false, "", err + } + if paddingLen > MaxPaddingLength { + return false, "", errors.ProtocolError{Message: "invalid padding length"} + } + if paddingLen > 0 { + _, err = io.CopyN(io.Discard, r, int64(paddingLen)) + if err != nil { + return false, "", err + } + } + return status[0] == 0, string(msgBuf), nil +} + +func WriteTCPResponse(w io.Writer, ok bool, msg string) error { + padding := tcpResponsePadding.String() + paddingLen := len(padding) + msgLen := len(msg) + sz := 1 + int(quicvarint.Len(uint64(msgLen))) + msgLen + + int(quicvarint.Len(uint64(paddingLen))) + paddingLen + buf := make([]byte, sz) + if ok { + buf[0] = 0 + } else { + buf[0] = 1 + } + i := varintPut(buf[1:], uint64(msgLen)) + i += copy(buf[1+i:], msg) + i += varintPut(buf[1+i:], uint64(paddingLen)) + copy(buf[1+i:], padding) + _, err := w.Write(buf) + return err +} + +// UDPMessage format: +// Session ID (uint32 BE) +// Packet ID (uint16 BE) +// Fragment ID (uint8) +// Fragment count (uint8) +// Address length (QUIC varint) +// Address (bytes) +// Data... + +type UDPMessage struct { + SessionID uint32 // 4 + PacketID uint16 // 2 + FragID uint8 // 1 + FragCount uint8 // 1 + Addr string // varint + bytes + Data []byte +} + +func (m *UDPMessage) HeaderSize() int { + lAddr := len(m.Addr) + return 4 + 2 + 1 + 1 + int(quicvarint.Len(uint64(lAddr))) + lAddr +} + +func (m *UDPMessage) Size() int { + return m.HeaderSize() + len(m.Data) +} + +func (m *UDPMessage) Serialize(buf []byte) int { + // Make sure the buffer is big enough + if len(buf) < m.Size() { + return -1 + } + binary.BigEndian.PutUint32(buf, m.SessionID) + binary.BigEndian.PutUint16(buf[4:], m.PacketID) + buf[6] = m.FragID + buf[7] = m.FragCount + i := varintPut(buf[8:], uint64(len(m.Addr))) + i += copy(buf[8+i:], m.Addr) + i += copy(buf[8+i:], m.Data) + return 8 + i +} + +func ParseUDPMessage(msg []byte) (*UDPMessage, error) { + m := &UDPMessage{} + buf := bytes.NewBuffer(msg) + if err := binary.Read(buf, binary.BigEndian, &m.SessionID); err != nil { + return nil, err + } + if err := binary.Read(buf, binary.BigEndian, &m.PacketID); err != nil { + return nil, err + } + if err := binary.Read(buf, binary.BigEndian, &m.FragID); err != nil { + return nil, err + } + if err := binary.Read(buf, binary.BigEndian, &m.FragCount); err != nil { + return nil, err + } + lAddr, err := quicvarint.Read(buf) + if err != nil { + return nil, err + } + if lAddr == 0 || lAddr > MaxMessageLength { + return nil, errors.ProtocolError{Message: "invalid address length"} + } + bs := buf.Bytes() + if len(bs) <= int(lAddr) { + // We use <= instead of < here as we expect at least one byte of data after the address + return nil, errors.ProtocolError{Message: "invalid message length"} + } + m.Addr = string(bs[:lAddr]) + m.Data = bs[lAddr:] + return m, nil +} + +// varintPut is like quicvarint.Append, but instead of appending to a slice, +// it writes to a fixed-size buffer. Returns the number of bytes written. +func varintPut(b []byte, i uint64) int { + if i <= maxVarInt1 { + b[0] = uint8(i) + return 1 + } + if i <= maxVarInt2 { + b[0] = uint8(i>>8) | 0x40 + b[1] = uint8(i) + return 2 + } + if i <= maxVarInt4 { + b[0] = uint8(i>>24) | 0x80 + b[1] = uint8(i >> 16) + b[2] = uint8(i >> 8) + b[3] = uint8(i) + return 4 + } + if i <= maxVarInt8 { + b[0] = uint8(i>>56) | 0xc0 + b[1] = uint8(i >> 48) + b[2] = uint8(i >> 40) + b[3] = uint8(i >> 32) + b[4] = uint8(i >> 24) + b[5] = uint8(i >> 16) + b[6] = uint8(i >> 8) + b[7] = uint8(i) + return 8 + } + panic(fmt.Sprintf("%#x doesn't fit into 62 bits", i)) +} diff --git a/third_party/hysteria-core/internal/protocol/proxy_test.go b/third_party/hysteria-core/internal/protocol/proxy_test.go new file mode 100644 index 0000000..9c16724 --- /dev/null +++ b/third_party/hysteria-core/internal/protocol/proxy_test.go @@ -0,0 +1,330 @@ +package protocol + +import ( + "bytes" + "reflect" + "strings" + "testing" +) + +func TestUDPMessage(t *testing.T) { + t.Run("buffer too small", func(t *testing.T) { + // Make sure Serialize returns -1 when the buffer is too small. + tBuf := make([]byte, 20) + if (&UDPMessage{ + SessionID: 66, + PacketID: 77, + FragID: 2, + FragCount: 5, + Addr: "random_addr", + Data: []byte("random_data"), + }).Serialize(tBuf) != -1 { + t.Error("Serialize() did not return -1 when the buffer was too small") + } + }) + + type fields struct { + SessionID uint32 + PacketID uint16 + FragID uint8 + FragCount uint8 + Addr string + Data []byte + } + tests := []struct { + name string + fields fields + want []byte + }{ + { + name: "test 1", + fields: fields{ + SessionID: 1, + PacketID: 1, + FragID: 0, + FragCount: 1, + Addr: "example.com:80", + Data: []byte("GET /nothing HTTP/1.1\r\n"), + }, + want: []byte{0x0, 0x0, 0x0, 0x1, 0x0, 0x1, 0x0, 0x1, 0xe, 0x65, 0x78, 0x61, 0x6d, 0x70, 0x6c, 0x65, 0x2e, 0x63, 0x6f, 0x6d, 0x3a, 0x38, 0x30, 0x47, 0x45, 0x54, 0x20, 0x2f, 0x6e, 0x6f, 0x74, 0x68, 0x69, 0x6e, 0x67, 0x20, 0x48, 0x54, 0x54, 0x50, 0x2f, 0x31, 0x2e, 0x31, 0xd, 0xa}, + }, + { + name: "test 2", + fields: fields{ + SessionID: 1329655244, + Addr: "some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long_some_random_goofy_ahh_address_which_is_very_long:9000", + PacketID: 62233, + FragID: 8, + FragCount: 19, + Data: []byte("God is great, beer is good, and people are crazy."), + }, + want: []byte{0x4f, 0x40, 0xed, 0xcc, 0xf3, 0x19, 0x8, 0x13, 0x41, 0xee, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x5f, 0x73, 0x6f, 0x6d, 0x65, 0x5f, 0x72, 0x61, 0x6e, 0x64, 0x6f, 0x6d, 0x5f, 0x67, 0x6f, 0x6f, 0x66, 0x79, 0x5f, 0x61, 0x68, 0x68, 0x5f, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x5f, 0x77, 0x68, 0x69, 0x63, 0x68, 0x5f, 0x69, 0x73, 0x5f, 0x76, 0x65, 0x72, 0x79, 0x5f, 0x6c, 0x6f, 0x6e, 0x67, 0x3a, 0x39, 0x30, 0x30, 0x30, 0x47, 0x6f, 0x64, 0x20, 0x69, 0x73, 0x20, 0x67, 0x72, 0x65, 0x61, 0x74, 0x2c, 0x20, 0x62, 0x65, 0x65, 0x72, 0x20, 0x69, 0x73, 0x20, 0x67, 0x6f, 0x6f, 0x64, 0x2c, 0x20, 0x61, 0x6e, 0x64, 0x20, 0x70, 0x65, 0x6f, 0x70, 0x6c, 0x65, 0x20, 0x61, 0x72, 0x65, 0x20, 0x63, 0x72, 0x61, 0x7a, 0x79, 0x2e}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + m := &UDPMessage{ + SessionID: tt.fields.SessionID, + Addr: tt.fields.Addr, + PacketID: tt.fields.PacketID, + FragID: tt.fields.FragID, + FragCount: tt.fields.FragCount, + Data: tt.fields.Data, + } + // Serialize + buf := make([]byte, MaxUDPSize) + n := m.Serialize(buf) + if got := buf[:n]; !reflect.DeepEqual(got, tt.want) { + t.Errorf("Serialize() = %v, want %v", got, tt.want) + } + // Parse back + if m2, err := ParseUDPMessage(tt.want); err != nil { + t.Errorf("ParseUDPMessage() error = %v", err) + } else { + if !reflect.DeepEqual(m2, m) { + t.Errorf("ParseUDPMessage() = %v, want %v", m2, m) + } + } + }) + } +} + +func TestMaximumUDPMessageFitsSerializationBuffer(t *testing.T) { + message := &UDPMessage{ + SessionID: 1, + FragCount: 1, + Addr: string(make([]byte, MaxAddressLength)), + Data: make([]byte, MaxUDPSize), + } + buf := make([]byte, MaxUDPMessageSize) + if got := message.Serialize(buf); got != message.Size() { + t.Fatalf("Serialize returned %d, want %d", got, message.Size()) + } +} + +// TestUDPMessageMalformed is to make sure ParseUDPMessage() fails (but not panic) on malformed data. +func TestUDPMessageMalformed(t *testing.T) { + tests := []struct { + name string + data []byte + }{ + { + name: "empty", + data: []byte{}, + }, + { + name: "zeroes 1", + data: []byte{0, 0, 0, 0}, + }, + { + name: "zeroes 2", + data: []byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, + }, + { + name: "incomplete 1", + data: []byte{0x66, 0xCC, 0xFF, 0xFF, 0x11, 0x22, 0x33, 0x44, 0x55}, + }, + { + name: "incomplete 2", + data: []byte{0x66, 0xCC, 0xFF, 0xFF, 0x11, 0x22, 0x33, 0x44, 0x90, 0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if _, err := ParseUDPMessage(tt.data); err == nil { + t.Errorf("ParseUDPMessage() should fail") + } + }) + } +} + +func TestReadTCPRequest(t *testing.T) { + tests := []struct { + name string + data []byte + want string + wantErr bool + }{ + { + name: "normal no padding", + data: []byte("\x0egoogle.com:443\x00"), + want: "google.com:443", + wantErr: false, + }, + { + name: "normal with padding", + data: []byte("\x0bholy.cc:443\x02gg"), + want: "holy.cc:443", + wantErr: false, + }, + { + name: "incomplete 1", + data: []byte("\x0bhoho"), + want: "", + wantErr: true, + }, + { + name: "incomplete 2", + data: []byte("\x0bholy.cc:443\x05x"), + want: "", + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + r := bytes.NewReader(tt.data) + got, err := ReadTCPRequest(r) + if (err != nil) != tt.wantErr { + t.Errorf("ReadTCPRequest() error = %v, wantErr %v", err, tt.wantErr) + return + } + if got != tt.want { + t.Errorf("ReadTCPRequest() got = %v, want %v", got, tt.want) + } + }) + } +} + +func TestWriteTCPRequest(t *testing.T) { + tests := []struct { + name string + addr string + wantW string // Just a prefix, we don't care about the padding + wantErr bool + }{ + { + name: "normal 1", + addr: "google.com:443", + wantW: "\x44\x01\x0egoogle.com:443", + wantErr: false, + }, + { + name: "normal 2", + addr: "client-api.arkoselabs.com:8080", + wantW: "\x44\x01\x1eclient-api.arkoselabs.com:8080", + wantErr: false, + }, + { + name: "empty", + addr: "", + wantW: "\x44\x01\x00", + wantErr: false, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + w := &bytes.Buffer{} + err := WriteTCPRequest(w, tt.addr) + if (err != nil) != tt.wantErr { + t.Errorf("WriteTCPRequest() error = %v, wantErr %v", err, tt.wantErr) + return + } + if gotW := w.String(); !(strings.HasPrefix(gotW, tt.wantW) && len(gotW) > len(tt.wantW)) { + t.Errorf("WriteTCPRequest() gotW = %v, want %v", gotW, tt.wantW) + } + }) + } +} + +func TestReadTCPResponse(t *testing.T) { + tests := []struct { + name string + data []byte + want bool + want1 string + wantErr bool + }{ + { + name: "normal ok no padding", + data: []byte("\x00\x0bhello world\x00"), + want: true, + want1: "hello world", + wantErr: false, + }, + { + name: "normal error with padding", + data: []byte("\x01\x06stop!!\x05xxxxx"), + want1: "stop!!", + wantErr: false, + }, + { + name: "normal ok no message with padding", + data: []byte("\x01\x00\x05xxxxx"), + want1: "", + wantErr: false, + }, + { + name: "incomplete 1", + data: []byte("\x00\x0bhoho"), + want1: "", + wantErr: true, + }, + { + name: "incomplete 2", + data: []byte("\x01\x05jesus\x05x"), + want1: "", + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + r := bytes.NewReader(tt.data) + got, got1, err := ReadTCPResponse(r) + if (err != nil) != tt.wantErr { + t.Errorf("ReadTCPResponse() error = %v, wantErr %v", err, tt.wantErr) + return + } + if got != tt.want { + t.Errorf("ReadTCPResponse() got = %v, want %v", got, tt.want) + } + if got1 != tt.want1 { + t.Errorf("ReadTCPResponse() got1 = %v, want %v", got1, tt.want1) + } + }) + } +} + +func TestWriteTCPResponse(t *testing.T) { + type args struct { + ok bool + msg string + } + tests := []struct { + name string + args args + wantW string // Just a prefix, we don't care about the padding + wantErr bool + }{ + { + name: "normal ok", + args: args{ok: true, msg: "hello world"}, + wantW: "\x00\x0bhello world", + wantErr: false, + }, + { + name: "normal error", + args: args{ok: false, msg: "stop!!"}, + wantW: "\x01\x06stop!!", + wantErr: false, + }, + { + name: "empty", + args: args{ok: true, msg: ""}, + wantW: "\x00\x00", + wantErr: false, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + w := &bytes.Buffer{} + err := WriteTCPResponse(w, tt.args.ok, tt.args.msg) + if (err != nil) != tt.wantErr { + t.Errorf("WriteTCPResponse() error = %v, wantErr %v", err, tt.wantErr) + return + } + if gotW := w.String(); !(strings.HasPrefix(gotW, tt.wantW) && len(gotW) > len(tt.wantW)) { + t.Errorf("WriteTCPResponse() gotW = %v, want %v", gotW, tt.wantW) + } + }) + } +} diff --git a/third_party/hysteria-core/internal/utils/atomic.go b/third_party/hysteria-core/internal/utils/atomic.go new file mode 100644 index 0000000..7739013 --- /dev/null +++ b/third_party/hysteria-core/internal/utils/atomic.go @@ -0,0 +1,54 @@ +package utils + +import ( + "sync/atomic" + "time" +) + +type AtomicTime struct { + v atomic.Value +} + +func NewAtomicTime(t time.Time) *AtomicTime { + a := &AtomicTime{} + a.Set(t) + return a +} + +func (t *AtomicTime) Set(new time.Time) { + t.v.Store(new) +} + +func (t *AtomicTime) Get() time.Time { + return t.v.Load().(time.Time) +} + +type Atomic[T any] struct { + v atomic.Value +} + +func (a *Atomic[T]) Load() T { + value := a.v.Load() + if value == nil { + var zero T + return zero + } + return value.(T) +} + +func (a *Atomic[T]) Store(value T) { + a.v.Store(value) +} + +func (a *Atomic[T]) Swap(new T) T { + old := a.v.Swap(new) + if old == nil { + var zero T + return zero + } + return old.(T) +} + +func (a *Atomic[T]) CompareAndSwap(old, new T) bool { + return a.v.CompareAndSwap(old, new) +} diff --git a/third_party/hysteria-core/internal/utils/qstream.go b/third_party/hysteria-core/internal/utils/qstream.go new file mode 100644 index 0000000..76519b6 --- /dev/null +++ b/third_party/hysteria-core/internal/utils/qstream.go @@ -0,0 +1,62 @@ +package utils + +import ( + "context" + "time" + + "github.com/apernet/quic-go" +) + +// QStream is a wrapper of quic.Stream that handles Close() in a way that +// makes more sense to us. By default, quic.Stream's Close() only closes +// the write side of the stream, not the read side. And if there is unread +// data, the stream is not really considered closed until either the data +// is drained or CancelRead() is called. +// References: +// - https://github.com/libp2p/go-libp2p/blob/master/p2p/transport/quic/stream.go +// - https://github.com/quic-go/quic-go/issues/3558 +// - https://github.com/quic-go/quic-go/issues/1599 +type QStream struct { + Stream *quic.Stream +} + +func (s *QStream) StreamID() quic.StreamID { + return s.Stream.StreamID() +} + +func (s *QStream) Read(p []byte) (n int, err error) { + return s.Stream.Read(p) +} + +func (s *QStream) CancelRead(code quic.StreamErrorCode) { + s.Stream.CancelRead(code) +} + +func (s *QStream) SetReadDeadline(t time.Time) error { + return s.Stream.SetReadDeadline(t) +} + +func (s *QStream) Write(p []byte) (n int, err error) { + return s.Stream.Write(p) +} + +func (s *QStream) Close() error { + s.Stream.CancelRead(0) + return s.Stream.Close() +} + +func (s *QStream) CancelWrite(code quic.StreamErrorCode) { + s.Stream.CancelWrite(code) +} + +func (s *QStream) Context() context.Context { + return s.Stream.Context() +} + +func (s *QStream) SetWriteDeadline(t time.Time) error { + return s.Stream.SetWriteDeadline(t) +} + +func (s *QStream) SetDeadline(t time.Time) error { + return s.Stream.SetDeadline(t) +} diff --git a/third_party/hysteria-core/server/.mockery.yaml b/third_party/hysteria-core/server/.mockery.yaml new file mode 100644 index 0000000..d73136a --- /dev/null +++ b/third_party/hysteria-core/server/.mockery.yaml @@ -0,0 +1,15 @@ +with-expecter: true +inpackage: true +dir: . +packages: + github.com/apernet/hysteria/core/v2/server: + interfaces: + udpIO: + config: + mockname: mockUDPIO + udpEventLogger: + config: + mockname: mockUDPEventLogger + UDPConn: + config: + mockname: mockUDPConn diff --git a/third_party/hysteria-core/server/config.go b/third_party/hysteria-core/server/config.go new file mode 100644 index 0000000..5ef5b5a --- /dev/null +++ b/third_party/hysteria-core/server/config.go @@ -0,0 +1,404 @@ +package server + +import ( + "crypto/tls" + "crypto/x509" + "io" + "net" + "net/http" + "sync/atomic" + "time" + + "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/congestion" + "github.com/apernet/hysteria/core/v2/internal/pmtud" + "github.com/apernet/hysteria/core/v2/internal/utils" + "github.com/apernet/quic-go" +) + +const ( + defaultStreamReceiveWindow = 8388608 // 8MB + defaultConnReceiveWindow = defaultStreamReceiveWindow * 5 / 2 // 20MB + defaultMaxIdleTimeout = 30 * time.Second + defaultMaxIncomingStreams = 1024 + defaultMaxIncomingUniStreams = 8 + defaultMaxHTTPHeaderBytes = 16 << 10 + defaultUDPIdleTimeout = 60 * time.Second + defaultMaxConnections = 512 + defaultMaxClientConnections = 32 + defaultMaxTCPHandlers = 1024 + defaultMaxClientTCPHandlers = 128 + defaultTCPRequestTimeout = 10 * time.Second + defaultAuthenticationTimeout = 10 * time.Second + defaultMaxUDPSessions = 1024 + defaultMaxClientUDPSessions = 64 +) + +type Config struct { + TLSConfig TLSConfig + QUICConfig QUICConfig + Conn net.PacketConn + StatelessResetKey *quic.StatelessResetKey + Cleanup io.Closer + RequestHook RequestHook + Outbound Outbound + CongestionConfig CongestionConfig + BandwidthConfig BandwidthConfig + IgnoreClientBandwidth bool + DisableUDP bool + UDPIdleTimeout time.Duration + // Resource limits are AutoCAR's security hardening over the v2.12.1 + // core. They bound work before an outbound socket exists, including + // unauthenticated QUIC connections and incomplete UDP fragments. + MaxConnections int + MaxClientConnections int + MaxTCPHandlers int + MaxClientTCPHandlers int + TCPRequestTimeout time.Duration + AuthenticationTimeout time.Duration + MaxHTTPHeaderBytes int + MaxUDPSessions int + MaxClientUDPSessions int + Authenticator Authenticator + EventLogger EventLogger + TrafficLogger TrafficLogger + MasqHandler http.Handler +} + +// fill fills the fields that are not set by the user with default values when possible, +// and returns an error if the user has not set a required field, or if a field is invalid. +func (c *Config) fill() error { + if len(c.TLSConfig.Certificates) == 0 && c.TLSConfig.GetCertificate == nil { + return errors.ConfigError{Field: "TLSConfig", Reason: "must set at least one of Certificates or GetCertificate"} + } + if c.QUICConfig.InitialStreamReceiveWindow == 0 { + c.QUICConfig.InitialStreamReceiveWindow = defaultStreamReceiveWindow + } else if c.QUICConfig.InitialStreamReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.InitialStreamReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.MaxStreamReceiveWindow == 0 { + c.QUICConfig.MaxStreamReceiveWindow = defaultStreamReceiveWindow + } else if c.QUICConfig.MaxStreamReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.MaxStreamReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.InitialConnectionReceiveWindow == 0 { + c.QUICConfig.InitialConnectionReceiveWindow = defaultConnReceiveWindow + } else if c.QUICConfig.InitialConnectionReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.InitialConnectionReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.MaxConnectionReceiveWindow == 0 { + c.QUICConfig.MaxConnectionReceiveWindow = defaultConnReceiveWindow + } else if c.QUICConfig.MaxConnectionReceiveWindow < 16384 { + return errors.ConfigError{Field: "QUICConfig.MaxConnectionReceiveWindow", Reason: "must be at least 16384"} + } + if c.QUICConfig.MaxIdleTimeout == 0 { + c.QUICConfig.MaxIdleTimeout = defaultMaxIdleTimeout + } else if c.QUICConfig.MaxIdleTimeout < 4*time.Second || c.QUICConfig.MaxIdleTimeout > 120*time.Second { + return errors.ConfigError{Field: "QUICConfig.MaxIdleTimeout", Reason: "must be between 4s and 120s"} + } + if c.QUICConfig.MaxIncomingStreams == 0 { + c.QUICConfig.MaxIncomingStreams = defaultMaxIncomingStreams + } else if c.QUICConfig.MaxIncomingStreams < 8 { + return errors.ConfigError{Field: "QUICConfig.MaxIncomingStreams", Reason: "must be at least 8"} + } + if c.QUICConfig.MaxIncomingUniStreams == 0 { + c.QUICConfig.MaxIncomingUniStreams = defaultMaxIncomingUniStreams + } else if c.QUICConfig.MaxIncomingUniStreams < 3 { + return errors.ConfigError{Field: "QUICConfig.MaxIncomingUniStreams", Reason: "must be at least 3"} + } + c.QUICConfig.DisablePathMTUDiscovery = c.QUICConfig.DisablePathMTUDiscovery || pmtud.DisablePathMTUDiscovery + var err error + c.CongestionConfig.Type, err = congestion.NormalizeType(c.CongestionConfig.Type) + if err != nil { + return errors.ConfigError{Field: "CongestionConfig.Type", Reason: err.Error()} + } + if c.CongestionConfig.Type == congestion.TypeBBR { + c.CongestionConfig.BBRProfile, err = congestion.NormalizeBBRProfile(c.CongestionConfig.BBRProfile) + if err != nil { + return errors.ConfigError{Field: "CongestionConfig.BBRProfile", Reason: err.Error()} + } + } + if c.Conn == nil { + return errors.ConfigError{Field: "Conn", Reason: "must be set"} + } + if c.Outbound == nil { + c.Outbound = &defaultOutbound{} + } + if c.BandwidthConfig.MaxTx != 0 && c.BandwidthConfig.MaxTx < 65536 { + return errors.ConfigError{Field: "BandwidthConfig.MaxTx", Reason: "must be at least 65536"} + } + if c.BandwidthConfig.MaxRx != 0 && c.BandwidthConfig.MaxRx < 65536 { + return errors.ConfigError{Field: "BandwidthConfig.MaxRx", Reason: "must be at least 65536"} + } + if c.UDPIdleTimeout == 0 { + c.UDPIdleTimeout = defaultUDPIdleTimeout + } else if c.UDPIdleTimeout < 2*time.Second || c.UDPIdleTimeout > 600*time.Second { + return errors.ConfigError{Field: "UDPIdleTimeout", Reason: "must be between 2s and 600s"} + } + if c.MaxConnections == 0 { + c.MaxConnections = defaultMaxConnections + } else if c.MaxConnections < 1 { + return errors.ConfigError{Field: "MaxConnections", Reason: "must be positive"} + } + if c.MaxClientConnections == 0 { + c.MaxClientConnections = min(defaultMaxClientConnections, c.MaxConnections) + } else if c.MaxClientConnections < 1 { + return errors.ConfigError{Field: "MaxClientConnections", Reason: "must be positive"} + } + if c.MaxClientConnections > c.MaxConnections { + return errors.ConfigError{Field: "MaxClientConnections", Reason: "must not exceed MaxConnections"} + } + if c.MaxTCPHandlers == 0 { + c.MaxTCPHandlers = defaultMaxTCPHandlers + } else if c.MaxTCPHandlers < 1 { + return errors.ConfigError{Field: "MaxTCPHandlers", Reason: "must be positive"} + } + if c.MaxClientTCPHandlers == 0 { + c.MaxClientTCPHandlers = min(defaultMaxClientTCPHandlers, c.MaxTCPHandlers) + } else if c.MaxClientTCPHandlers < 1 { + return errors.ConfigError{Field: "MaxClientTCPHandlers", Reason: "must be positive"} + } + if c.MaxClientTCPHandlers > c.MaxTCPHandlers { + return errors.ConfigError{Field: "MaxClientTCPHandlers", Reason: "must not exceed MaxTCPHandlers"} + } + if c.TCPRequestTimeout == 0 { + c.TCPRequestTimeout = defaultTCPRequestTimeout + } else if c.TCPRequestTimeout < time.Second || c.TCPRequestTimeout > 60*time.Second { + return errors.ConfigError{Field: "TCPRequestTimeout", Reason: "must be between 1s and 60s"} + } + if c.AuthenticationTimeout == 0 { + c.AuthenticationTimeout = defaultAuthenticationTimeout + } else if c.AuthenticationTimeout < time.Second || c.AuthenticationTimeout > 60*time.Second { + return errors.ConfigError{Field: "AuthenticationTimeout", Reason: "must be between 1s and 60s"} + } + if c.MaxHTTPHeaderBytes == 0 { + c.MaxHTTPHeaderBytes = defaultMaxHTTPHeaderBytes + } else if c.MaxHTTPHeaderBytes < 8<<10 || c.MaxHTTPHeaderBytes > 1<<20 { + return errors.ConfigError{Field: "MaxHTTPHeaderBytes", Reason: "must be between 8192 and 1048576"} + } + if c.MaxUDPSessions == 0 { + c.MaxUDPSessions = defaultMaxUDPSessions + } else if c.MaxUDPSessions < 1 { + return errors.ConfigError{Field: "MaxUDPSessions", Reason: "must be positive"} + } + if c.MaxClientUDPSessions == 0 { + c.MaxClientUDPSessions = min(defaultMaxClientUDPSessions, c.MaxUDPSessions) + } else if c.MaxClientUDPSessions < 1 { + return errors.ConfigError{Field: "MaxClientUDPSessions", Reason: "must be positive"} + } + if c.MaxClientUDPSessions > c.MaxUDPSessions { + return errors.ConfigError{Field: "MaxClientUDPSessions", Reason: "must not exceed MaxUDPSessions"} + } + if c.Authenticator == nil { + return errors.ConfigError{Field: "Authenticator", Reason: "must be set"} + } + return nil +} + +// TLSConfig contains the TLS configuration fields that we want to expose to the user. +type TLSConfig struct { + Certificates []tls.Certificate + GetCertificate func(info *tls.ClientHelloInfo) (*tls.Certificate, error) + ClientCAs *x509.CertPool + ECHKeys []tls.EncryptedClientHelloKey + GetECHKeys func(info *tls.ClientHelloInfo) ([]tls.EncryptedClientHelloKey, error) +} + +// QUICConfig contains the QUIC configuration fields that we want to expose to the user. +type QUICConfig struct { + InitialStreamReceiveWindow uint64 + MaxStreamReceiveWindow uint64 + InitialConnectionReceiveWindow uint64 + MaxConnectionReceiveWindow uint64 + MaxIdleTimeout time.Duration + MaxIncomingStreams int64 + MaxIncomingUniStreams int64 + DisablePathMTUDiscovery bool // The server may still override this to true on unsupported platforms. + DisableGSO bool +} + +type CongestionConfig struct { + Type string + BBRProfile string +} + +// RequestHook allows filtering and modifying requests before the server connects to the remote. +// A request will only be hooked if Check returns true. +// The returned byte slice, if not empty, will be sent to the remote before proxying - this is +// mainly for "putting back" the content read from the client for sniffing, etc. +// Return a non-nil error to abort the connection. +// Note that due to the current architectural limitations, it can only inspect the first packet +// of a UDP connection. It also cannot put back any data as the first packet is always sent as-is. +type RequestHook interface { + Check(isUDP bool, reqAddr string) bool + TCP(stream HyStream, reqAddr *string) ([]byte, error) + UDP(data []byte, reqAddr *string) error +} + +// Outbound provides the implementation of how the server should connect to remote servers. +// Although UDP includes a reqAddr, the implementation does not necessarily have to use it +// to make a "connected" UDP connection that does not accept packets from other addresses. +// In fact, the default implementation simply uses net.ListenUDP for a "full-cone" behavior. +// CheckUDP is used to check if a UDP packet to reqAddr is permitted (useful for e.g. ACL). +type Outbound interface { + TCP(reqAddr string) (net.Conn, error) + UDP(reqAddr string) (UDPConn, error) + CheckUDP(reqAddr string) error +} + +// UDPConn is like net.PacketConn, but uses string for addresses. +type UDPConn interface { + ReadFrom(b []byte) (int, string, error) + WriteTo(b []byte, addr string) (int, error) + Close() error +} + +type defaultOutbound struct{} + +var defaultOutboundDialer = net.Dialer{ + Timeout: 10 * time.Second, +} + +func (o *defaultOutbound) TCP(reqAddr string) (net.Conn, error) { + return defaultOutboundDialer.Dial("tcp", reqAddr) +} + +func (o *defaultOutbound) UDP(reqAddr string) (UDPConn, error) { + conn, err := net.ListenUDP("udp", nil) + if err != nil { + return nil, err + } + return &defaultUDPConn{conn}, nil +} + +func (o *defaultOutbound) CheckUDP(reqAddr string) error { + return nil +} + +type defaultUDPConn struct { + *net.UDPConn +} + +func (c *defaultUDPConn) ReadFrom(b []byte) (int, string, error) { + n, addr, err := c.UDPConn.ReadFrom(b) + if addr != nil { + return n, addr.String(), err + } else { + return n, "", err + } +} + +func (c *defaultUDPConn) WriteTo(b []byte, addr string) (int, error) { + uAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return 0, err + } + return c.UDPConn.WriteTo(b, uAddr) +} + +// BandwidthConfig describes the maximum bandwidth that the server can use, in bytes per second. +type BandwidthConfig struct { + MaxTx uint64 + MaxRx uint64 + DisableLossCompensation bool +} + +// Authenticator is an interface that provides authentication logic. +type Authenticator interface { + Authenticate(addr net.Addr, auth string, tx uint64) (ok bool, id string) +} + +// EventLogger is an interface that provides logging logic. +type EventLogger interface { + Connect(addr net.Addr, id string, tx uint64) + Disconnect(addr net.Addr, id string, err error) + TCPRequest(addr net.Addr, id, reqAddr string) + TCPError(addr net.Addr, id, reqAddr string, err error) + UDPRequest(addr net.Addr, id string, sessionID uint32, reqAddr string) + UDPError(addr net.Addr, id string, sessionID uint32, err error) +} + +type HyStream interface { + StreamID() quic.StreamID + Read(p []byte) (n int, err error) + Write(p []byte) (n int, err error) + Close() error + SetReadDeadline(t time.Time) error + SetWriteDeadline(t time.Time) error + SetDeadline(t time.Time) error +} + +// TrafficLogger is an interface that provides traffic logging logic. +// Tx/Rx in this context refers to the server-remote (proxy target) perspective. +// Tx is the bytes sent from the server to the remote. +// Rx is the bytes received by the server from the remote. +// Apart from logging, the Log function can also return false to signal +// that the client should be disconnected. This can be used to implement +// bandwidth limits or post-connection authentication, for example. +// The implementation of this interface must be thread-safe. +type TrafficLogger interface { + LogTraffic(id string, tx, rx uint64) (ok bool) + LogOnlineState(id string, online bool) + TraceStream(stream HyStream, stats *StreamStats) + UntraceStream(stream HyStream) +} + +type StreamState int + +const ( + // StreamStateInitial indicates the initial state of a stream. + // Client has opened the stream, but we have not received the proxy request yet. + StreamStateInitial StreamState = iota + + // StreamStateHooking indicates that the hook (usually sniff) is processing. + // Client has sent the proxy request, but sniff requires more data to complete. + StreamStateHooking + + // StreamStateConnecting indicates that we are connecting to the proxy target. + StreamStateConnecting + + // StreamStateEstablished indicates the proxy is established. + StreamStateEstablished + + // StreamStateClosed indicates the stream is closed. + StreamStateClosed +) + +func (s StreamState) String() string { + switch s { + case StreamStateInitial: + return "init" + case StreamStateHooking: + return "hook" + case StreamStateConnecting: + return "connect" + case StreamStateEstablished: + return "estab" + case StreamStateClosed: + return "closed" + default: + return "unknown" + } +} + +type StreamStats struct { + State utils.Atomic[StreamState] + + AuthID string + ConnID uint32 + InitialTime time.Time + + ReqAddr utils.Atomic[string] + HookedReqAddr utils.Atomic[string] + + Tx atomic.Uint64 + Rx atomic.Uint64 + + LastActiveTime utils.Atomic[time.Time] +} + +func (s *StreamStats) setHookedReqAddr(addr string) { + if addr != s.ReqAddr.Load() { + s.HookedReqAddr.Store(addr) + } +} diff --git a/third_party/hysteria-core/server/copy.go b/third_party/hysteria-core/server/copy.go new file mode 100644 index 0000000..ea916d8 --- /dev/null +++ b/third_party/hysteria-core/server/copy.go @@ -0,0 +1,80 @@ +package server + +import ( + "errors" + "io" + "sync" + "time" +) + +var errDisconnect = errors.New("traffic logger requested disconnect") + +var copyBufPool = sync.Pool{ + New: func() any { + b := make([]byte, 32*1024) + return &b + }, +} + +func copyBufferLog(dst io.Writer, src io.Reader, log func(n uint64) bool) error { + bufp := copyBufPool.Get().(*[]byte) + buf := *bufp + defer copyBufPool.Put(bufp) + + for { + nr, er := src.Read(buf) + if nr > 0 { + if !log(uint64(nr)) { + // Log returns false, which means that the client should be disconnected + return errDisconnect + } + _, ew := dst.Write(buf[0:nr]) + if ew != nil { + return ew + } + } + if er != nil { + if er == io.EOF { + // EOF should not be considered as an error + return nil + } + return er + } + } +} + +func copyTwoWayEx(id string, serverRw, remoteRw io.ReadWriter, l TrafficLogger, stats *StreamStats) error { + errChan := make(chan error, 2) + go func() { + errChan <- copyBufferLog(serverRw, remoteRw, func(n uint64) bool { + stats.LastActiveTime.Store(time.Now()) + stats.Rx.Add(n) + return l.LogTraffic(id, 0, n) + }) + }() + go func() { + errChan <- copyBufferLog(remoteRw, serverRw, func(n uint64) bool { + stats.LastActiveTime.Store(time.Now()) + stats.Tx.Add(n) + return l.LogTraffic(id, n, 0) + }) + }() + // Block until one of the two goroutines returns + return <-errChan +} + +// copyTwoWay is the "fast-path" version of copyTwoWayEx that does not log traffic or update stream stats. +// It uses the built-in io.Copy instead of our own copyBufferLog. +func copyTwoWay(serverRw, remoteRw io.ReadWriter) error { + errChan := make(chan error, 2) + go func() { + _, err := io.Copy(serverRw, remoteRw) + errChan <- err + }() + go func() { + _, err := io.Copy(remoteRw, serverRw) + errChan <- err + }() + // Block until one of the two goroutines returns + return <-errChan +} diff --git a/third_party/hysteria-core/server/copy_benchmark_test.go b/third_party/hysteria-core/server/copy_benchmark_test.go new file mode 100644 index 0000000..0f17bea --- /dev/null +++ b/third_party/hysteria-core/server/copy_benchmark_test.go @@ -0,0 +1,23 @@ +package server + +import ( + "bytes" + "io" + "testing" +) + +func BenchmarkCopyBufferLog(b *testing.B) { + srcData := make([]byte, 1024*1024) // 1MB + for i := range srcData { + srcData[i] = byte(i) + } + + b.ReportAllocs() + b.ResetTimer() + + for i := 0; i < b.N; i++ { + src := bytes.NewReader(srcData) + dst := io.Discard + copyBufferLog(dst, src, func(n uint64) bool { return true }) + } +} diff --git a/third_party/hysteria-core/server/mock_UDPConn.go b/third_party/hysteria-core/server/mock_UDPConn.go new file mode 100644 index 0000000..5f3d2e9 --- /dev/null +++ b/third_party/hysteria-core/server/mock_UDPConn.go @@ -0,0 +1,197 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package server + +import mock "github.com/stretchr/testify/mock" + +// mockUDPConn is an autogenerated mock type for the UDPConn type +type mockUDPConn struct { + mock.Mock +} + +type mockUDPConn_Expecter struct { + mock *mock.Mock +} + +func (_m *mockUDPConn) EXPECT() *mockUDPConn_Expecter { + return &mockUDPConn_Expecter{mock: &_m.Mock} +} + +// Close provides a mock function with no fields +func (_m *mockUDPConn) Close() error { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for Close") + } + + var r0 error + if rf, ok := ret.Get(0).(func() error); ok { + r0 = rf() + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockUDPConn_Close_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Close' +type mockUDPConn_Close_Call struct { + *mock.Call +} + +// Close is a helper method to define mock.On call +func (_e *mockUDPConn_Expecter) Close() *mockUDPConn_Close_Call { + return &mockUDPConn_Close_Call{Call: _e.mock.On("Close")} +} + +func (_c *mockUDPConn_Close_Call) Run(run func()) *mockUDPConn_Close_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *mockUDPConn_Close_Call) Return(_a0 error) *mockUDPConn_Close_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockUDPConn_Close_Call) RunAndReturn(run func() error) *mockUDPConn_Close_Call { + _c.Call.Return(run) + return _c +} + +// ReadFrom provides a mock function with given fields: b +func (_m *mockUDPConn) ReadFrom(b []byte) (int, string, error) { + ret := _m.Called(b) + + if len(ret) == 0 { + panic("no return value specified for ReadFrom") + } + + var r0 int + var r1 string + var r2 error + if rf, ok := ret.Get(0).(func([]byte) (int, string, error)); ok { + return rf(b) + } + if rf, ok := ret.Get(0).(func([]byte) int); ok { + r0 = rf(b) + } else { + r0 = ret.Get(0).(int) + } + + if rf, ok := ret.Get(1).(func([]byte) string); ok { + r1 = rf(b) + } else { + r1 = ret.Get(1).(string) + } + + if rf, ok := ret.Get(2).(func([]byte) error); ok { + r2 = rf(b) + } else { + r2 = ret.Error(2) + } + + return r0, r1, r2 +} + +// mockUDPConn_ReadFrom_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ReadFrom' +type mockUDPConn_ReadFrom_Call struct { + *mock.Call +} + +// ReadFrom is a helper method to define mock.On call +// - b []byte +func (_e *mockUDPConn_Expecter) ReadFrom(b interface{}) *mockUDPConn_ReadFrom_Call { + return &mockUDPConn_ReadFrom_Call{Call: _e.mock.On("ReadFrom", b)} +} + +func (_c *mockUDPConn_ReadFrom_Call) Run(run func(b []byte)) *mockUDPConn_ReadFrom_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte)) + }) + return _c +} + +func (_c *mockUDPConn_ReadFrom_Call) Return(_a0 int, _a1 string, _a2 error) *mockUDPConn_ReadFrom_Call { + _c.Call.Return(_a0, _a1, _a2) + return _c +} + +func (_c *mockUDPConn_ReadFrom_Call) RunAndReturn(run func([]byte) (int, string, error)) *mockUDPConn_ReadFrom_Call { + _c.Call.Return(run) + return _c +} + +// WriteTo provides a mock function with given fields: b, addr +func (_m *mockUDPConn) WriteTo(b []byte, addr string) (int, error) { + ret := _m.Called(b, addr) + + if len(ret) == 0 { + panic("no return value specified for WriteTo") + } + + var r0 int + var r1 error + if rf, ok := ret.Get(0).(func([]byte, string) (int, error)); ok { + return rf(b, addr) + } + if rf, ok := ret.Get(0).(func([]byte, string) int); ok { + r0 = rf(b, addr) + } else { + r0 = ret.Get(0).(int) + } + + if rf, ok := ret.Get(1).(func([]byte, string) error); ok { + r1 = rf(b, addr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// mockUDPConn_WriteTo_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'WriteTo' +type mockUDPConn_WriteTo_Call struct { + *mock.Call +} + +// WriteTo is a helper method to define mock.On call +// - b []byte +// - addr string +func (_e *mockUDPConn_Expecter) WriteTo(b interface{}, addr interface{}) *mockUDPConn_WriteTo_Call { + return &mockUDPConn_WriteTo_Call{Call: _e.mock.On("WriteTo", b, addr)} +} + +func (_c *mockUDPConn_WriteTo_Call) Run(run func(b []byte, addr string)) *mockUDPConn_WriteTo_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte), args[1].(string)) + }) + return _c +} + +func (_c *mockUDPConn_WriteTo_Call) Return(_a0 int, _a1 error) *mockUDPConn_WriteTo_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *mockUDPConn_WriteTo_Call) RunAndReturn(run func([]byte, string) (int, error)) *mockUDPConn_WriteTo_Call { + _c.Call.Return(run) + return _c +} + +// newMockUDPConn creates a new instance of mockUDPConn. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func newMockUDPConn(t interface { + mock.TestingT + Cleanup(func()) +}) *mockUDPConn { + mock := &mockUDPConn{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/server/mock_udpEventLogger.go b/third_party/hysteria-core/server/mock_udpEventLogger.go new file mode 100644 index 0000000..e1d3db9 --- /dev/null +++ b/third_party/hysteria-core/server/mock_udpEventLogger.go @@ -0,0 +1,100 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package server + +import mock "github.com/stretchr/testify/mock" + +// mockUDPEventLogger is an autogenerated mock type for the udpEventLogger type +type mockUDPEventLogger struct { + mock.Mock +} + +type mockUDPEventLogger_Expecter struct { + mock *mock.Mock +} + +func (_m *mockUDPEventLogger) EXPECT() *mockUDPEventLogger_Expecter { + return &mockUDPEventLogger_Expecter{mock: &_m.Mock} +} + +// Close provides a mock function with given fields: sessionID, err +func (_m *mockUDPEventLogger) Close(sessionID uint32, err error) { + _m.Called(sessionID, err) +} + +// mockUDPEventLogger_Close_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Close' +type mockUDPEventLogger_Close_Call struct { + *mock.Call +} + +// Close is a helper method to define mock.On call +// - sessionID uint32 +// - err error +func (_e *mockUDPEventLogger_Expecter) Close(sessionID interface{}, err interface{}) *mockUDPEventLogger_Close_Call { + return &mockUDPEventLogger_Close_Call{Call: _e.mock.On("Close", sessionID, err)} +} + +func (_c *mockUDPEventLogger_Close_Call) Run(run func(sessionID uint32, err error)) *mockUDPEventLogger_Close_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint32), args[1].(error)) + }) + return _c +} + +func (_c *mockUDPEventLogger_Close_Call) Return() *mockUDPEventLogger_Close_Call { + _c.Call.Return() + return _c +} + +func (_c *mockUDPEventLogger_Close_Call) RunAndReturn(run func(uint32, error)) *mockUDPEventLogger_Close_Call { + _c.Run(run) + return _c +} + +// New provides a mock function with given fields: sessionID, reqAddr +func (_m *mockUDPEventLogger) New(sessionID uint32, reqAddr string) { + _m.Called(sessionID, reqAddr) +} + +// mockUDPEventLogger_New_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'New' +type mockUDPEventLogger_New_Call struct { + *mock.Call +} + +// New is a helper method to define mock.On call +// - sessionID uint32 +// - reqAddr string +func (_e *mockUDPEventLogger_Expecter) New(sessionID interface{}, reqAddr interface{}) *mockUDPEventLogger_New_Call { + return &mockUDPEventLogger_New_Call{Call: _e.mock.On("New", sessionID, reqAddr)} +} + +func (_c *mockUDPEventLogger_New_Call) Run(run func(sessionID uint32, reqAddr string)) *mockUDPEventLogger_New_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint32), args[1].(string)) + }) + return _c +} + +func (_c *mockUDPEventLogger_New_Call) Return() *mockUDPEventLogger_New_Call { + _c.Call.Return() + return _c +} + +func (_c *mockUDPEventLogger_New_Call) RunAndReturn(run func(uint32, string)) *mockUDPEventLogger_New_Call { + _c.Run(run) + return _c +} + +// newMockUDPEventLogger creates a new instance of mockUDPEventLogger. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func newMockUDPEventLogger(t interface { + mock.TestingT + Cleanup(func()) +}) *mockUDPEventLogger { + mock := &mockUDPEventLogger{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/server/mock_udpIO.go b/third_party/hysteria-core/server/mock_udpIO.go new file mode 100644 index 0000000..bb512c0 --- /dev/null +++ b/third_party/hysteria-core/server/mock_udpIO.go @@ -0,0 +1,290 @@ +// Code generated by mockery v2.53.5. DO NOT EDIT. + +package server + +import ( + protocol "github.com/apernet/hysteria/core/v2/internal/protocol" + mock "github.com/stretchr/testify/mock" +) + +// mockUDPIO is an autogenerated mock type for the udpIO type +type mockUDPIO struct { + mock.Mock +} + +type mockUDPIO_Expecter struct { + mock *mock.Mock +} + +func (_m *mockUDPIO) EXPECT() *mockUDPIO_Expecter { + return &mockUDPIO_Expecter{mock: &_m.Mock} +} + +// CheckUDP provides a mock function with given fields: reqAddr +func (_m *mockUDPIO) CheckUDP(reqAddr string) error { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for CheckUDP") + } + + var r0 error + if rf, ok := ret.Get(0).(func(string) error); ok { + r0 = rf(reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockUDPIO_CheckUDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckUDP' +type mockUDPIO_CheckUDP_Call struct { + *mock.Call +} + +// CheckUDP is a helper method to define mock.On call +// - reqAddr string +func (_e *mockUDPIO_Expecter) CheckUDP(reqAddr interface{}) *mockUDPIO_CheckUDP_Call { + return &mockUDPIO_CheckUDP_Call{Call: _e.mock.On("CheckUDP", reqAddr)} +} + +func (_c *mockUDPIO_CheckUDP_Call) Run(run func(reqAddr string)) *mockUDPIO_CheckUDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *mockUDPIO_CheckUDP_Call) Return(_a0 error) *mockUDPIO_CheckUDP_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockUDPIO_CheckUDP_Call) RunAndReturn(run func(string) error) *mockUDPIO_CheckUDP_Call { + _c.Call.Return(run) + return _c +} + +// Hook provides a mock function with given fields: data, reqAddr +func (_m *mockUDPIO) Hook(data []byte, reqAddr *string) error { + ret := _m.Called(data, reqAddr) + + if len(ret) == 0 { + panic("no return value specified for Hook") + } + + var r0 error + if rf, ok := ret.Get(0).(func([]byte, *string) error); ok { + r0 = rf(data, reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockUDPIO_Hook_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Hook' +type mockUDPIO_Hook_Call struct { + *mock.Call +} + +// Hook is a helper method to define mock.On call +// - data []byte +// - reqAddr *string +func (_e *mockUDPIO_Expecter) Hook(data interface{}, reqAddr interface{}) *mockUDPIO_Hook_Call { + return &mockUDPIO_Hook_Call{Call: _e.mock.On("Hook", data, reqAddr)} +} + +func (_c *mockUDPIO_Hook_Call) Run(run func(data []byte, reqAddr *string)) *mockUDPIO_Hook_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte), args[1].(*string)) + }) + return _c +} + +func (_c *mockUDPIO_Hook_Call) Return(_a0 error) *mockUDPIO_Hook_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockUDPIO_Hook_Call) RunAndReturn(run func([]byte, *string) error) *mockUDPIO_Hook_Call { + _c.Call.Return(run) + return _c +} + +// ReceiveMessage provides a mock function with no fields +func (_m *mockUDPIO) ReceiveMessage() (*protocol.UDPMessage, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for ReceiveMessage") + } + + var r0 *protocol.UDPMessage + var r1 error + if rf, ok := ret.Get(0).(func() (*protocol.UDPMessage, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() *protocol.UDPMessage); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*protocol.UDPMessage) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// mockUDPIO_ReceiveMessage_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ReceiveMessage' +type mockUDPIO_ReceiveMessage_Call struct { + *mock.Call +} + +// ReceiveMessage is a helper method to define mock.On call +func (_e *mockUDPIO_Expecter) ReceiveMessage() *mockUDPIO_ReceiveMessage_Call { + return &mockUDPIO_ReceiveMessage_Call{Call: _e.mock.On("ReceiveMessage")} +} + +func (_c *mockUDPIO_ReceiveMessage_Call) Run(run func()) *mockUDPIO_ReceiveMessage_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *mockUDPIO_ReceiveMessage_Call) Return(_a0 *protocol.UDPMessage, _a1 error) *mockUDPIO_ReceiveMessage_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *mockUDPIO_ReceiveMessage_Call) RunAndReturn(run func() (*protocol.UDPMessage, error)) *mockUDPIO_ReceiveMessage_Call { + _c.Call.Return(run) + return _c +} + +// SendMessage provides a mock function with given fields: _a0, _a1 +func (_m *mockUDPIO) SendMessage(_a0 []byte, _a1 *protocol.UDPMessage) error { + ret := _m.Called(_a0, _a1) + + if len(ret) == 0 { + panic("no return value specified for SendMessage") + } + + var r0 error + if rf, ok := ret.Get(0).(func([]byte, *protocol.UDPMessage) error); ok { + r0 = rf(_a0, _a1) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockUDPIO_SendMessage_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendMessage' +type mockUDPIO_SendMessage_Call struct { + *mock.Call +} + +// SendMessage is a helper method to define mock.On call +// - _a0 []byte +// - _a1 *protocol.UDPMessage +func (_e *mockUDPIO_Expecter) SendMessage(_a0 interface{}, _a1 interface{}) *mockUDPIO_SendMessage_Call { + return &mockUDPIO_SendMessage_Call{Call: _e.mock.On("SendMessage", _a0, _a1)} +} + +func (_c *mockUDPIO_SendMessage_Call) Run(run func(_a0 []byte, _a1 *protocol.UDPMessage)) *mockUDPIO_SendMessage_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]byte), args[1].(*protocol.UDPMessage)) + }) + return _c +} + +func (_c *mockUDPIO_SendMessage_Call) Return(_a0 error) *mockUDPIO_SendMessage_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockUDPIO_SendMessage_Call) RunAndReturn(run func([]byte, *protocol.UDPMessage) error) *mockUDPIO_SendMessage_Call { + _c.Call.Return(run) + return _c +} + +// UDP provides a mock function with given fields: reqAddr +func (_m *mockUDPIO) UDP(reqAddr string) (UDPConn, error) { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for UDP") + } + + var r0 UDPConn + var r1 error + if rf, ok := ret.Get(0).(func(string) (UDPConn, error)); ok { + return rf(reqAddr) + } + if rf, ok := ret.Get(0).(func(string) UDPConn); ok { + r0 = rf(reqAddr) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(UDPConn) + } + } + + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(reqAddr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// mockUDPIO_UDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UDP' +type mockUDPIO_UDP_Call struct { + *mock.Call +} + +// UDP is a helper method to define mock.On call +// - reqAddr string +func (_e *mockUDPIO_Expecter) UDP(reqAddr interface{}) *mockUDPIO_UDP_Call { + return &mockUDPIO_UDP_Call{Call: _e.mock.On("UDP", reqAddr)} +} + +func (_c *mockUDPIO_UDP_Call) Run(run func(reqAddr string)) *mockUDPIO_UDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *mockUDPIO_UDP_Call) Return(_a0 UDPConn, _a1 error) *mockUDPIO_UDP_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *mockUDPIO_UDP_Call) RunAndReturn(run func(string) (UDPConn, error)) *mockUDPIO_UDP_Call { + _c.Call.Return(run) + return _c +} + +// newMockUDPIO creates a new instance of mockUDPIO. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func newMockUDPIO(t interface { + mock.TestingT + Cleanup(func()) +}) *mockUDPIO { + mock := &mockUDPIO{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/third_party/hysteria-core/server/resource_limits_test.go b/third_party/hysteria-core/server/resource_limits_test.go new file mode 100644 index 0000000..0219725 --- /dev/null +++ b/third_party/hysteria-core/server/resource_limits_test.go @@ -0,0 +1,409 @@ +package server + +import ( + "context" + "crypto/tls" + "errors" + "net" + "sync" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/apernet/hysteria/core/v2/internal/frag" + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +type testAuthenticator struct{} + +func (testAuthenticator) Authenticate(net.Addr, string, uint64) (bool, string) { + return true, "test" +} + +func TestPerSourceDefaultsRespectSmallGlobalLimits(t *testing.T) { + packetConn, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer packetConn.Close() + config := &Config{ + TLSConfig: TLSConfig{Certificates: []tls.Certificate{{}}}, + Conn: packetConn, + MaxConnections: 2, + MaxTCPHandlers: 3, + MaxUDPSessions: 4, + Authenticator: testAuthenticator{}, + } + if err := config.fill(); err != nil { + t.Fatal(err) + } + if config.MaxClientConnections != 2 || config.MaxClientTCPHandlers != 3 || config.MaxClientUDPSessions != 4 { + t.Fatalf("per-source defaults = (%d, %d, %d), want (2, 3, 4)", config.MaxClientConnections, config.MaxClientTCPHandlers, config.MaxClientUDPSessions) + } +} + +func TestConnectionAdmissionCoversHandshakeLifecycle(t *testing.T) { + slots := make(chan struct{}, 1) + clients := newKeyedLimiter(1) + firstContext, cancelFirst := context.WithCancel(context.Background()) + defer cancelFirst() + if _, err := admitConnection(firstContext, slots, clients, "192.0.2.10"); err != nil { + t.Fatalf("first connection admission: %v", err) + } + if _, err := admitConnection(context.Background(), slots, clients, "192.0.2.11"); !errors.Is(err, errConnectionCapacity) { + t.Fatalf("second connection admission error = %v, want capacity", err) + } + cancelFirst() + deadline := time.Now().Add(time.Second) + for len(slots) != 0 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if len(slots) != 0 { + t.Fatal("connection slot was not released when its context ended") + } + secondContext, cancelSecond := context.WithCancel(context.Background()) + defer cancelSecond() + if _, err := admitConnection(secondContext, slots, clients, "192.0.2.11"); err != nil { + t.Fatalf("released connection slot was not reusable: %v", err) + } +} + +func TestConnectionAdmissionLimitIsSharedAcrossSourceConnections(t *testing.T) { + global := make(chan struct{}, 3) + clients := newKeyedLimiter(1) + firstContext, cancelFirst := context.WithCancel(context.Background()) + defer cancelFirst() + if _, err := admitConnection(firstContext, global, clients, "192.0.2.10"); err != nil { + t.Fatal(err) + } + if _, err := admitConnection(context.Background(), global, clients, "192.0.2.10"); !errors.Is(err, errConnectionCapacity) { + t.Fatalf("same-source connection error = %v, want capacity", err) + } + otherContext, cancelOther := context.WithCancel(context.Background()) + defer cancelOther() + if _, err := admitConnection(otherContext, global, clients, "192.0.2.11"); err != nil { + t.Fatalf("different source was rejected: %v", err) + } + if len(global) != 2 || clients.count("192.0.2.10") != 1 || clients.count("192.0.2.11") != 1 { + t.Fatalf("unexpected admission state: global=%d first=%d other=%d", len(global), clients.count("192.0.2.10"), clients.count("192.0.2.11")) + } + + cancelFirst() + cancelOther() + deadline := time.Now().Add(time.Second) + for (len(global) != 0 || clients.count("192.0.2.10") != 0 || clients.count("192.0.2.11") != 0) && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if len(global) != 0 || clients.count("192.0.2.10") != 0 || clients.count("192.0.2.11") != 0 { + t.Fatal("connection admission retained capacity or empty source entries") + } +} + +func TestSourceIPKeyGroupsIPv6PrefixAndIgnoresPort(t *testing.T) { + first := &net.UDPAddr{IP: net.ParseIP("2001:db8:1234:5678::1"), Port: 443} + second := &net.UDPAddr{IP: net.ParseIP("2001:db8:1234:5678:ffff::2"), Port: 8443} + other := &net.UDPAddr{IP: net.ParseIP("2001:db8:1234:5679::1"), Port: 443} + if sourceIPKey(first) != sourceIPKey(second) { + t.Fatalf("same IPv6 /64 produced different keys: %q and %q", sourceIPKey(first), sourceIPKey(second)) + } + if sourceIPKey(first) == sourceIPKey(other) { + t.Fatalf("different IPv6 /64 prefixes shared key %q", sourceIPKey(first)) + } + if got := sourceIPKey(&net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 443}); got != "192.0.2.10" { + t.Fatalf("IPv4 key = %q, want address without port", got) + } +} + +func TestUDPSessionAdmissionPrecedesDefragmentation(t *testing.T) { + ioMock := newMockUDPIO(t) + events := newMockUDPEventLogger(t) + global := make(chan struct{}, 1) + first := newUDPSessionManager(ioMock, events, time.Minute, 1, global, "client-a", nil) + second := newUDPSessionManager(ioMock, events, time.Minute, 1, global, "client-b", nil) + + incomplete := func(id uint32) *protocol.UDPMessage { + return &protocol.UDPMessage{ + SessionID: id, + PacketID: 1, + FragID: 0, + FragCount: 2, + Addr: "example.test:443", + Data: []byte("partial"), + } + } + + first.feed(incomplete(1)) + second.feed(incomplete(2)) + if got := first.Count(); got != 1 { + t.Fatalf("first manager count = %d, want 1", got) + } + if got := second.Count(); got != 0 { + t.Fatalf("global admission allowed second incomplete session: %d", got) + } + if got := len(global); got != 1 { + t.Fatalf("global slots = %d, want 1", got) + } + + events.EXPECT().Close(uint32(1), nil).Once() + first.cleanup(false) + second.feed(incomplete(2)) + if got := second.Count(); got != 1 { + t.Fatalf("released slot was not reusable: count = %d", got) + } + events.EXPECT().Close(uint32(2), nil).Once() + second.cleanup(false) +} + +func TestUDPSessionRejectsExcessiveFragmentsWithoutState(t *testing.T) { + ioMock := newMockUDPIO(t) + events := newMockUDPEventLogger(t) + global := make(chan struct{}, 1) + manager := newUDPSessionManager(ioMock, events, time.Minute, 1, global, "client", nil) + manager.feed(&protocol.UDPMessage{ + SessionID: 7, + PacketID: 1, + FragID: 0, + FragCount: frag.MaxFragments + 1, + Addr: "example.test:443", + Data: []byte("partial"), + }) + if got := manager.Count(); got != 0 { + t.Fatalf("malformed fragment allocated %d sessions", got) + } + if got := len(global); got != 0 { + t.Fatalf("malformed fragment consumed %d global slots", got) + } +} + +func TestUDPSessionLimitIsSharedAcrossClientConnections(t *testing.T) { + ioMock := newMockUDPIO(t) + events := newMockUDPEventLogger(t) + global := make(chan struct{}, 3) + clients := newKeyedLimiter(1) + first := newUDPSessionManager(ioMock, events, time.Minute, 2, global, "192.0.2.10", clients) + second := newUDPSessionManager(ioMock, events, time.Minute, 2, global, "192.0.2.10", clients) + other := newUDPSessionManager(ioMock, events, time.Minute, 2, global, "192.0.2.11", clients) + message := func(id uint32) *protocol.UDPMessage { + return &protocol.UDPMessage{ + SessionID: id, + PacketID: 1, + FragID: 0, + FragCount: 2, + Addr: "example.test:443", + Data: []byte("partial"), + } + } + + first.feed(message(1)) + second.feed(message(2)) + other.feed(message(3)) + if first.Count() != 1 || second.Count() != 0 || other.Count() != 1 { + t.Fatalf("cross-connection counts = (%d, %d, %d), want (1, 0, 1)", first.Count(), second.Count(), other.Count()) + } + if got := clients.count("192.0.2.10"); got != 1 { + t.Fatalf("shared client count = %d, want 1", got) + } + + events.EXPECT().Close(uint32(1), nil).Once() + first.cleanup(false) + second.feed(message(2)) + if second.Count() != 1 { + t.Fatal("released per-client slot was not reusable by another connection") + } + events.EXPECT().Close(uint32(2), nil).Once() + events.EXPECT().Close(uint32(3), nil).Once() + second.cleanup(false) + other.cleanup(false) + if clients.count("192.0.2.10") != 0 || clients.count("192.0.2.11") != 0 { + t.Fatal("per-client limiter retained empty source entries") + } +} + +func TestPendingTCPHeaderTimesOutAndReleasesSlot(t *testing.T) { + stream := newDeadlineStream() + slots := make(chan struct{}, 1) + handler := &h3sHandler{ + config: &Config{TCPRequestTimeout: 20 * time.Millisecond}, + tcpSlots: slots, + } + release, admitted := handler.admitStream(stream) + if !admitted || release == nil { + t.Fatal("stream was not admitted") + } + if got := len(slots); got != 1 { + t.Fatalf("TCP handler slots after admission = %d, want 1", got) + } + done := make(chan struct{}) + go func() { + handler.handleTCPRequest(stream, "test-auth") + close(done) + }() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("pending TCP header did not time out") + } + if !stream.wasClosed() { + t.Fatal("timed-out stream was not closed") + } + release() + release() + if got := len(slots); got != 0 { + t.Fatalf("TCP handler slot was not released exactly once: %d", got) + } +} + +func TestTCPHandlerLimitIsSharedAcrossClientConnections(t *testing.T) { + global := make(chan struct{}, 3) + clients := newKeyedLimiter(1) + first := &h3sHandler{ + config: &Config{TCPRequestTimeout: time.Second}, + clientKey: "192.0.2.10", + tcpSlots: global, + tcpClientSlots: clients, + } + second := &h3sHandler{ + config: &Config{TCPRequestTimeout: time.Second}, + clientKey: "192.0.2.10", + tcpSlots: global, + tcpClientSlots: clients, + } + other := &h3sHandler{ + config: &Config{TCPRequestTimeout: time.Second}, + clientKey: "192.0.2.11", + tcpSlots: global, + tcpClientSlots: clients, + } + + firstRelease, admitted := first.admitStream(newDeadlineStream()) + if !admitted { + t.Fatal("first source stream was rejected") + } + if release, admitted := second.admitStream(newDeadlineStream()); admitted || release != nil { + t.Fatal("second connection bypassed the shared per-source TCP limit") + } + otherRelease, admitted := other.admitStream(newDeadlineStream()) + if !admitted { + t.Fatal("different source was rejected while global capacity remained") + } + if len(global) != 2 || clients.count("192.0.2.10") != 1 || clients.count("192.0.2.11") != 1 { + t.Fatalf("unexpected limiter state: global=%d first=%d other=%d", len(global), clients.count("192.0.2.10"), clients.count("192.0.2.11")) + } + + firstRelease() + secondRelease, admitted := second.admitStream(newDeadlineStream()) + if !admitted { + t.Fatal("released per-source TCP slot was not reusable across connections") + } + secondRelease() + otherRelease() + if len(global) != 0 || clients.count("192.0.2.10") != 0 || clients.count("192.0.2.11") != 0 { + t.Fatal("TCP limiters retained capacity or empty source entries") + } +} + +func TestAuthStateWaitsForAuthenticationUpdate(t *testing.T) { + handler := &h3sHandler{} + handler.authMutex.Lock() + result := make(chan struct { + id string + ok bool + }, 1) + go func() { + id, ok := handler.authState() + result <- struct { + id string + ok bool + }{id: id, ok: ok} + }() + + select { + case <-result: + t.Fatal("authState returned during an in-progress authentication update") + case <-time.After(20 * time.Millisecond): + } + handler.authID = "authenticated-user" + handler.authenticated = true + handler.authMutex.Unlock() + + select { + case state := <-result: + if !state.ok || state.id != "authenticated-user" { + t.Fatalf("authState = (%q, %v), want authenticated user", state.id, state.ok) + } + case <-time.After(time.Second): + t.Fatal("authState remained blocked after authentication completed") + } +} + +func TestAuthenticationExpiryIsOneShotAndFailClosed(t *testing.T) { + handler := &h3sHandler{} + if !handler.expireAuthentication() { + t.Fatal("first unauthenticated expiry was ignored") + } + if handler.expireAuthentication() { + t.Fatal("authentication expiry fired more than once") + } + if id, ok := handler.authState(); ok || id != "" { + t.Fatalf("expired auth state = (%q, %v), want unauthenticated", id, ok) + } + + authenticated := &h3sHandler{authenticated: true, authID: "client"} + if authenticated.expireAuthentication() { + t.Fatal("authenticated session was expired") + } +} + +type deadlineStream struct { + mu sync.Mutex + deadline time.Time + closed bool +} + +func newDeadlineStream() *deadlineStream { return &deadlineStream{} } + +func (s *deadlineStream) StreamID() quic.StreamID { return 0 } +func (s *deadlineStream) Write(p []byte) (int, error) { return len(p), nil } +func (s *deadlineStream) SetWriteDeadline(time.Time) error { return nil } +func (s *deadlineStream) SetDeadline(deadline time.Time) error { + return s.SetReadDeadline(deadline) +} +func (s *deadlineStream) SetReadDeadline(deadline time.Time) error { + s.mu.Lock() + s.deadline = deadline + s.mu.Unlock() + return nil +} +func (s *deadlineStream) Read([]byte) (int, error) { + s.mu.Lock() + deadline := s.deadline + s.mu.Unlock() + if deadline.IsZero() { + return 0, errors.New("missing read deadline") + } + if delay := time.Until(deadline); delay > 0 { + time.Sleep(delay) + } + return 0, timeoutError{} +} +func (s *deadlineStream) Close() error { + s.mu.Lock() + s.closed = true + s.mu.Unlock() + return nil +} +func (s *deadlineStream) wasClosed() bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.closed +} + +type timeoutError struct{} + +func (timeoutError) Error() string { return "deadline exceeded" } +func (timeoutError) Timeout() bool { return true } +func (timeoutError) Temporary() bool { return true } + +var _ HyStream = (*deadlineStream)(nil) diff --git a/third_party/hysteria-core/server/server.go b/third_party/hysteria-core/server/server.go new file mode 100644 index 0000000..c9b0abe --- /dev/null +++ b/third_party/hysteria-core/server/server.go @@ -0,0 +1,601 @@ +package server + +import ( + "context" + crand "crypto/rand" + "crypto/tls" + "errors" + "math/rand" + "net" + "net/http" + "net/netip" + "sync" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/quicvarint" + + coreErrs "github.com/apernet/hysteria/core/v2/errors" + "github.com/apernet/hysteria/core/v2/internal/congestion" + "github.com/apernet/hysteria/core/v2/internal/protocol" + "github.com/apernet/hysteria/core/v2/internal/utils" +) + +const ( + closeErrCodeOK = 0x100 // HTTP3 ErrCodeNoError + closeErrCodeTrafficLimitReached = 0x107 // HTTP3 ErrCodeExcessiveLoad +) + +type Server interface { + Serve() error + Close() error +} + +func convertToStdTLSConfig(config *Config) *tls.Config { + var clientAuth tls.ClientAuthType + if config.TLSConfig.ClientCAs != nil { + clientAuth = tls.RequireAndVerifyClientCert + } else { + clientAuth = tls.NoClientCert + } + return http3.ConfigureTLSConfig(&tls.Config{ + Certificates: config.TLSConfig.Certificates, + GetCertificate: config.TLSConfig.GetCertificate, + ClientCAs: config.TLSConfig.ClientCAs, + ClientAuth: clientAuth, + EncryptedClientHelloKeys: config.TLSConfig.ECHKeys, + GetEncryptedClientHelloKeys: config.TLSConfig.GetECHKeys, + }) +} + +func NewServer(config *Config) (Server, error) { + if err := config.fill(); err != nil { + return nil, err + } + tlsConfig := convertToStdTLSConfig(config) + quicConfig := &quic.Config{ + InitialStreamReceiveWindow: config.QUICConfig.InitialStreamReceiveWindow, + MaxStreamReceiveWindow: config.QUICConfig.MaxStreamReceiveWindow, + InitialConnectionReceiveWindow: config.QUICConfig.InitialConnectionReceiveWindow, + MaxConnectionReceiveWindow: config.QUICConfig.MaxConnectionReceiveWindow, + MaxIdleTimeout: config.QUICConfig.MaxIdleTimeout, + MaxIncomingStreams: config.QUICConfig.MaxIncomingStreams, + MaxIncomingUniStreams: config.QUICConfig.MaxIncomingUniStreams, + DisablePathMTUDiscovery: config.QUICConfig.DisablePathMTUDiscovery, + EnableDatagrams: true, + MaxDatagramFrameSize: protocol.MaxDatagramFrameSize, + AssumePeerMaxDatagramFrameSize: protocol.MaxDatagramFrameSize, + DisablePathManager: true, + } + srk := config.StatelessResetKey + if srk == nil { + var k quic.StatelessResetKey + if _, err := crand.Read(k[:]); err != nil { + return nil, err + } + srk = &k + } + tr := &quic.Transport{ + Conn: config.Conn, + DisableGSO: config.QUICConfig.DisableGSO, + StatelessResetKey: srk, + // Always require QUIC Retry before allocating a connection. This proves + // return-path reachability and prevents spoofed Initial packets from + // creating handshake state. + VerifySourceAddress: func(net.Addr) bool { return true }, + } + connSlots := make(chan struct{}, config.MaxConnections) + connClientSlots := newKeyedLimiter(config.MaxClientConnections) + tr.ConnContext = func(ctx context.Context, info *quic.ClientInfo) (context.Context, error) { + return admitConnection(ctx, connSlots, connClientSlots, sourceIPKey(info.RemoteAddr)) + } + listener, err := tr.Listen(tlsConfig, quicConfig) + if err != nil { + err = errors.Join(err, tr.Close(), config.Conn.Close()) + if config.Cleanup != nil { + err = errors.Join(err, config.Cleanup.Close()) + } + return nil, err + } + return &serverImpl{ + config: config, + tr: tr, + listener: listener, + tcpSlots: make(chan struct{}, config.MaxTCPHandlers), + tcpClientSlots: newKeyedLimiter(config.MaxClientTCPHandlers), + udpSlots: make(chan struct{}, config.MaxUDPSessions), + udpClientSlots: newKeyedLimiter(config.MaxClientUDPSessions), + }, nil +} + +type serverImpl struct { + config *Config + tr *quic.Transport + listener *quic.Listener + tcpSlots chan struct{} + tcpClientSlots *keyedLimiter + udpSlots chan struct{} + udpClientSlots *keyedLimiter +} + +func (s *serverImpl) Serve() error { + for { + conn, err := s.listener.Accept(context.Background()) + if err != nil { + return err + } + go s.handleClient(conn) + } +} + +var errConnectionCapacity = errors.New("connection capacity reached") + +// admitConnection acquires capacity after Retry has validated the source but +// before quic-go allocates handshake state. The connection context is canceled +// on every handshake failure or established-connection close, which releases +// the slot for the complete lifecycle. +func admitConnection(ctx context.Context, slots chan struct{}, clientSlots *keyedLimiter, clientKey string) (context.Context, error) { + if !clientSlots.tryAcquire(clientKey) { + return nil, errConnectionCapacity + } + select { + case slots <- struct{}{}: + go func() { + <-ctx.Done() + <-slots + clientSlots.release(clientKey) + }() + return ctx, nil + default: + clientSlots.release(clientKey) + return nil, errConnectionCapacity + } +} + +func (s *serverImpl) Close() error { + err := errors.Join(s.listener.Close(), s.tr.Close(), s.config.Conn.Close()) + if s.config.Cleanup != nil { + err = errors.Join(err, s.config.Cleanup.Close()) + } + return err +} + +func (s *serverImpl) handleClient(conn *quic.Conn) { + handler := newH3sHandler(s.config, conn, s.tcpSlots, s.tcpClientSlots, s.udpSlots, s.udpClientSlots) + authTimer := time.AfterFunc(s.config.AuthenticationTimeout, func() { + if handler.expireAuthentication() { + _ = conn.CloseWithError(closeErrCodeOK, "authentication timeout") + } + }) + h3s := http3.Server{ + Handler: handler, + MaxHeaderBytes: s.config.MaxHTTPHeaderBytes, + StreamAdmission: handler.AdmitStream, + StreamDispatcher: handler.ProxyStreamHijacker, + } + err := h3s.ServeQUICConn(conn) + authTimer.Stop() + // If the client is authenticated, we need to log the disconnect event + if authID, authenticated := handler.authState(); authenticated { + if tl := s.config.TrafficLogger; tl != nil { + tl.LogOnlineState(authID, false) + } + if el := s.config.EventLogger; el != nil { + el.Disconnect(conn.RemoteAddr(), authID, err) + } + } + _ = conn.CloseWithError(closeErrCodeOK, "") +} + +type h3sHandler struct { + config *Config + conn *quic.Conn + + authenticated bool + authExpired bool + authMutex sync.RWMutex + authID string + connID uint32 // a random id for dump streams + clientKey string + tcpSlots chan struct{} + tcpClientSlots *keyedLimiter + udpSlots chan struct{} + udpClientSlots *keyedLimiter +} + +func newH3sHandler( + config *Config, + conn *quic.Conn, + tcpSlots chan struct{}, + tcpClientSlots *keyedLimiter, + udpSlots chan struct{}, + udpClientSlots *keyedLimiter, +) *h3sHandler { + return &h3sHandler{ + config: config, + conn: conn, + connID: rand.Uint32(), + clientKey: sourceIPKey(conn.RemoteAddr()), + tcpSlots: tcpSlots, + tcpClientSlots: tcpClientSlots, + udpSlots: udpSlots, + udpClientSlots: udpClientSlots, + } +} + +func (h *h3sHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost && r.Host == protocol.URLHost && r.URL.Path == protocol.URLPath { + h.authMutex.Lock() + defer h.authMutex.Unlock() + if h.authExpired { + h.masqHandler(w, r) + return + } + if h.authenticated { + // Already authenticated + protocol.AuthResponseToHeader(w.Header(), protocol.AuthResponse{ + UDPEnabled: !h.config.DisableUDP, + Rx: h.config.BandwidthConfig.MaxRx, + RxAuto: h.config.IgnoreClientBandwidth, + }) + w.WriteHeader(protocol.StatusAuthOK) + return + } + authReq := protocol.AuthRequestFromHeader(r.Header) + actualTx := authReq.Rx + ok, id := h.config.Authenticator.Authenticate(h.conn.RemoteAddr(), authReq.Auth, actualTx) + if ok { + // Set authenticated flag + h.authenticated = true + h.authID = id + if h.config.IgnoreClientBandwidth { + // Ignore client bandwidth and use the configured congestion controller. + congestion.UseConfigured(h.conn, h.config.CongestionConfig.Type, h.config.CongestionConfig.BBRProfile) + actualTx = 0 + } else { + // actualTx = min(serverTx, clientRx) + if h.config.BandwidthConfig.MaxTx > 0 && actualTx > h.config.BandwidthConfig.MaxTx { + // We have a maxTx limit and the client is asking for more than that, + // return and use the limit instead + actualTx = h.config.BandwidthConfig.MaxTx + } + if actualTx > 0 { + congestion.UseBrutal(h.conn, actualTx, h.config.BandwidthConfig.DisableLossCompensation) + } else { + // Client doesn't know its own bandwidth, use the configured congestion controller. + congestion.UseConfigured(h.conn, h.config.CongestionConfig.Type, h.config.CongestionConfig.BBRProfile) + } + } + // Auth OK, send response + protocol.AuthResponseToHeader(w.Header(), protocol.AuthResponse{ + UDPEnabled: !h.config.DisableUDP, + Rx: h.config.BandwidthConfig.MaxRx, + RxAuto: h.config.IgnoreClientBandwidth, + }) + w.WriteHeader(protocol.StatusAuthOK) + // Call event logger + if tl := h.config.TrafficLogger; tl != nil { + tl.LogOnlineState(id, true) + } + if el := h.config.EventLogger; el != nil { + el.Connect(h.conn.RemoteAddr(), id, actualTx) + } + // Initialize UDP session manager (if UDP is enabled) + // We use sync.Once to make sure that only one goroutine is started, + // as ServeHTTP may be called by multiple goroutines simultaneously + if !h.config.DisableUDP { + sm := newUDPSessionManager( + &udpIOImpl{h.conn, id, h.config.TrafficLogger, h.config.RequestHook, h.config.Outbound}, + &udpEventLoggerImpl{h.conn, id, h.config.EventLogger}, + h.config.UDPIdleTimeout, + h.config.MaxClientUDPSessions, + h.udpSlots, + h.clientKey, + h.udpClientSlots, + ) + go sm.Run() + } + } else { + // Auth failed, pretend to be a normal HTTP server + h.masqHandler(w, r) + } + } else { + // Not an auth request, pretend to be a normal HTTP server + h.masqHandler(w, r) + } +} + +func sourceIPKey(address net.Addr) string { + if udpAddress, ok := address.(*net.UDPAddr); ok { + if ip, valid := netip.AddrFromSlice(udpAddress.IP); valid { + return sourcePrefixKey(ip) + } + } + host, _, err := net.SplitHostPort(address.String()) + if err == nil { + if ip, parseErr := netip.ParseAddr(host); parseErr == nil { + return sourcePrefixKey(ip) + } + return host + } + return address.String() +} + +// sourcePrefixKey prevents a client with an ordinary IPv6 /64 from evading +// every per-source budget by rotating interface identifiers. IPv4 remains +// keyed by the individual address. The complete remote address is still used +// for logging; only resource admission uses this normalized key. +func sourcePrefixKey(address netip.Addr) string { + address = address.Unmap().WithZone("") + if address.Is4() { + return address.String() + } + return netip.PrefixFrom(address, 64).Masked().String() +} + +func (h *h3sHandler) ProxyStreamHijacker(ft http3.FrameType, stream *quic.Stream, err error) (bool, error) { + authID, authenticated := h.authState() + if err != nil || !authenticated { + return false, nil + } + + switch ft { + case protocol.FrameTypeTCPRequest: + // StreamDispatcher only peeks the frame type. Consume it so ReadTCPRequest + // starts at address length, matching pre-upgrade StreamHijacker behavior. + if _, err := quicvarint.Read(quicvarint.NewReader(stream)); err != nil { + return false, err + } + // Wraps the stream with QStream, which handles Close() properly + qStream := &utils.QStream{Stream: stream} + // Run synchronously in quic-go's per-stream worker. StreamAdmission's + // release callback is deferred by that worker, so the global TCP slot + // remains held for the complete proxy lifetime, not merely until the + // request frame has been dispatched. + h.handleTCPRequest(qStream, authID) + return true, nil + default: + return false, nil + } +} + +// authState synchronizes HTTP authentication with stream dispatch and +// disconnect accounting. In particular, a stream arriving concurrently with +// the authentication response cannot observe authenticated=true with a stale +// or empty authID. +func (h *h3sHandler) authState() (string, bool) { + h.authMutex.RLock() + defer h.authMutex.RUnlock() + return h.authID, h.authenticated +} + +// expireAuthentication atomically prevents late authentication. It returns +// true exactly once when an unauthenticated connection crosses its deadline. +func (h *h3sHandler) expireAuthentication() bool { + h.authMutex.Lock() + defer h.authMutex.Unlock() + if h.authenticated || h.authExpired { + return false + } + h.authExpired = true + return true +} + +// AdmitStream applies listener-wide capacity and a first-byte deadline before +// the HTTP/3 layer peeks a frame type. The local quic-go fork guarantees that +// release runs once after dispatch or ordinary HTTP handling completes. +func (h *h3sHandler) AdmitStream(stream *quic.Stream) (func(), bool) { + return h.admitStream(stream) +} + +type readDeadlineSetter interface { + SetReadDeadline(time.Time) error +} + +func (h *h3sHandler) admitStream(stream readDeadlineSetter) (func(), bool) { + if !h.tcpClientSlots.tryAcquire(h.clientKey) { + return nil, false + } + select { + case h.tcpSlots <- struct{}{}: + default: + h.tcpClientSlots.release(h.clientKey) + return nil, false + } + if err := stream.SetReadDeadline(time.Now().Add(h.config.TCPRequestTimeout)); err != nil { + <-h.tcpSlots + h.tcpClientSlots.release(h.clientKey) + return nil, false + } + var once sync.Once + return func() { + once.Do(func() { + <-h.tcpSlots + h.tcpClientSlots.release(h.clientKey) + }) + }, true +} + +func (h *h3sHandler) handleTCPRequest(stream HyStream, authID string) { + trafficLogger := h.config.TrafficLogger + streamStats := &StreamStats{ + AuthID: authID, + ConnID: h.connID, + InitialTime: time.Now(), + } + streamStats.State.Store(StreamStateInitial) + streamStats.LastActiveTime.Store(time.Now()) + defer func() { + streamStats.State.Store(StreamStateClosed) + }() + if trafficLogger != nil { + trafficLogger.TraceStream(stream, streamStats) + defer trafficLogger.UntraceStream(stream) + } + + // Read request + _ = stream.SetReadDeadline(time.Now().Add(h.config.TCPRequestTimeout)) + reqAddr, err := protocol.ReadTCPRequest(stream) + if err != nil { + _ = stream.Close() + return + } + _ = stream.SetReadDeadline(time.Time{}) + streamStats.ReqAddr.Store(reqAddr) + // Call the hook if set + var putback []byte + var hooked bool + if h.config.RequestHook != nil { + hooked = h.config.RequestHook.Check(false, reqAddr) + // When the hook is enabled, the server should always accept a connection + // so that the client will send whatever request the hook wants to see. + // This is essentially a server-side fast-open. + if hooked { + streamStats.State.Store(StreamStateHooking) + _ = protocol.WriteTCPResponse(stream, true, "RequestHook enabled") + putback, err = h.config.RequestHook.TCP(stream, &reqAddr) + if err != nil { + _ = stream.Close() + return + } + streamStats.setHookedReqAddr(reqAddr) + } + } + // Log the event + if h.config.EventLogger != nil { + h.config.EventLogger.TCPRequest(h.conn.RemoteAddr(), authID, reqAddr) + } + // Dial target + streamStats.State.Store(StreamStateConnecting) + tConn, err := h.config.Outbound.TCP(reqAddr) + if err != nil { + if !hooked { + _ = protocol.WriteTCPResponse(stream, false, err.Error()) + } + _ = stream.Close() + // Log the error + if h.config.EventLogger != nil { + h.config.EventLogger.TCPError(h.conn.RemoteAddr(), authID, reqAddr, err) + } + return + } + if !hooked { + _ = protocol.WriteTCPResponse(stream, true, "Connected") + } + streamStats.State.Store(StreamStateEstablished) + // Put back the data if the hook requested + if len(putback) > 0 { + n, _ := tConn.Write(putback) + streamStats.Tx.Add(uint64(n)) + } + // Start proxying + if trafficLogger != nil { + err = copyTwoWayEx(authID, stream, tConn, trafficLogger, streamStats) + } else { + // Use the fast path if no traffic logger is set + err = copyTwoWay(stream, tConn) + } + if h.config.EventLogger != nil { + h.config.EventLogger.TCPError(h.conn.RemoteAddr(), authID, reqAddr, err) + } + // Cleanup + _ = tConn.Close() + _ = stream.Close() + // Disconnect the client if TrafficLogger requested + if err == errDisconnect { + _ = h.conn.CloseWithError(closeErrCodeTrafficLimitReached, "") + } +} + +func (h *h3sHandler) masqHandler(w http.ResponseWriter, r *http.Request) { + if h.config.MasqHandler != nil { + h.config.MasqHandler.ServeHTTP(w, r) + } else { + // Return 404 for everything + http.NotFound(w, r) + } +} + +// udpIOImpl is the IO implementation for udpSessionManager with TrafficLogger support +type udpIOImpl struct { + Conn *quic.Conn + AuthID string + TrafficLogger TrafficLogger + RequestHook RequestHook + Outbound Outbound +} + +func (io *udpIOImpl) ReceiveMessage() (*protocol.UDPMessage, error) { + for { + msg, err := io.Conn.ReceiveDatagram(context.Background()) + if err != nil { + // Connection error, this will stop the session manager + return nil, err + } + udpMsg, err := protocol.ParseUDPMessage(msg) + if err != nil { + // Invalid message, this is fine - just wait for the next + continue + } + if io.TrafficLogger != nil { + ok := io.TrafficLogger.LogTraffic(io.AuthID, uint64(len(udpMsg.Data)), 0) + if !ok { + // TrafficLogger requested to disconnect the client + _ = io.Conn.CloseWithError(closeErrCodeTrafficLimitReached, "") + return nil, errDisconnect + } + } + return udpMsg, nil + } +} + +func (io *udpIOImpl) SendMessage(buf []byte, msg *protocol.UDPMessage) error { + if io.TrafficLogger != nil { + ok := io.TrafficLogger.LogTraffic(io.AuthID, 0, uint64(len(msg.Data))) + if !ok { + // TrafficLogger requested to disconnect the client + _ = io.Conn.CloseWithError(closeErrCodeTrafficLimitReached, "") + return errDisconnect + } + } + msgN := msg.Serialize(buf) + if msgN < 0 { + return coreErrs.ProtocolError{Message: "UDP message exceeds serialization limit"} + } + return io.Conn.SendDatagram(buf[:msgN]) +} + +func (io *udpIOImpl) Hook(data []byte, reqAddr *string) error { + if io.RequestHook != nil && io.RequestHook.Check(true, *reqAddr) { + return io.RequestHook.UDP(data, reqAddr) + } else { + return nil + } +} + +func (io *udpIOImpl) UDP(reqAddr string) (UDPConn, error) { + return io.Outbound.UDP(reqAddr) +} + +func (io *udpIOImpl) CheckUDP(reqAddr string) error { + return io.Outbound.CheckUDP(reqAddr) +} + +type udpEventLoggerImpl struct { + Conn *quic.Conn + AuthID string + EventLogger EventLogger +} + +func (l *udpEventLoggerImpl) New(sessionID uint32, reqAddr string) { + if l.EventLogger != nil { + l.EventLogger.UDPRequest(l.Conn.RemoteAddr(), l.AuthID, sessionID, reqAddr) + } +} + +func (l *udpEventLoggerImpl) Close(sessionID uint32, err error) { + if l.EventLogger != nil { + l.EventLogger.UDPError(l.Conn.RemoteAddr(), l.AuthID, sessionID, err) + } +} diff --git a/third_party/hysteria-core/server/udp.go b/third_party/hysteria-core/server/udp.go new file mode 100644 index 0000000..d13a81d --- /dev/null +++ b/third_party/hysteria-core/server/udp.go @@ -0,0 +1,470 @@ +package server + +import ( + "errors" + "math/rand" + "sync" + "time" + + "github.com/apernet/quic-go" + + "github.com/apernet/hysteria/core/v2/internal/frag" + "github.com/apernet/hysteria/core/v2/internal/protocol" + "github.com/apernet/hysteria/core/v2/internal/utils" +) + +const ( + idleCleanupInterval = 1 * time.Second + maxSessionACLCache = 256 +) + +type udpIO interface { + ReceiveMessage() (*protocol.UDPMessage, error) + SendMessage([]byte, *protocol.UDPMessage) error + Hook(data []byte, reqAddr *string) error + UDP(reqAddr string) (UDPConn, error) + CheckUDP(reqAddr string) error +} + +type udpEventLogger interface { + New(sessionID uint32, reqAddr string) + Close(sessionID uint32, err error) +} + +type udpSessionEntry struct { + ID uint32 + OverrideAddr string // Ignore the address in the UDP message, always use this if not empty + OriginalAddr string // The original address in the UDP message + D *frag.Defragger + Last *utils.AtomicTime + IO udpIO + + DialFunc func(addr string, firstMsgData []byte) (conn UDPConn, actualAddr string, err error) + ExitFunc func(err error) + + conn UDPConn + connLock sync.Mutex + closed bool + + aclCache map[string]error +} + +func newUDPSessionEntry( + id uint32, io udpIO, + dialFunc func(string, []byte) (UDPConn, string, error), + exitFunc func(error), +) (e *udpSessionEntry) { + e = &udpSessionEntry{ + ID: id, + D: &frag.Defragger{}, + Last: utils.NewAtomicTime(time.Now()), + IO: io, + + DialFunc: dialFunc, + ExitFunc: exitFunc, + } + + return e +} + +// CloseWithErr closes the session and calls ExitFunc with the given error. +// A nil error indicates the session is cleaned up due to timeout. +func (e *udpSessionEntry) CloseWithErr(err error) { + // We need this lock to ensure not to create conn after session exit + e.connLock.Lock() + + if e.closed { + // Already closed + e.connLock.Unlock() + return + } + + e.closed = true + if e.conn != nil { + _ = e.conn.Close() + } + e.connLock.Unlock() + + e.ExitFunc(err) +} + +// Feed feeds a UDP message to the session. +// If the message itself is a complete message, or it completes a fragmented message, +// the message is written to the session's UDP connection, and the number of bytes +// written is returned. +// Otherwise, 0 and nil are returned. +func (e *udpSessionEntry) Feed(msg *protocol.UDPMessage) (int, error) { + e.Last.Set(time.Now()) + dfMsg := e.D.Feed(msg) + if dfMsg == nil { + return 0, nil + } + + if e.conn == nil { + err := e.initConn(dfMsg) + if err != nil { + return 0, err + } + if e.OverrideAddr == "" { + e.aclCache = map[string]error{dfMsg.Addr: nil} + } + } + + addr := dfMsg.Addr + if e.OverrideAddr != "" { + addr = e.OverrideAddr + } else if err := e.checkAddr(addr); err != nil { + return 0, err + } + + return e.conn.WriteTo(dfMsg.Data, addr) +} + +// checkAddr checks outbound policy for the given address. +// The decision is cached in e.aclCache for future use. +func (e *udpSessionEntry) checkAddr(addr string) error { + if decision, ok := e.aclCache[addr]; ok { + return decision + } + decision := e.IO.CheckUDP(addr) + if len(e.aclCache) >= maxSessionACLCache { + for k := range e.aclCache { + delete(e.aclCache, k) + break + } + } + if e.aclCache == nil { + e.aclCache = make(map[string]error, 4) + } + e.aclCache[addr] = decision + return decision +} + +// initConn initializes the UDP connection of the session. +// If no error is returned, the e.conn is set to the new connection. +func (e *udpSessionEntry) initConn(firstMsg *protocol.UDPMessage) error { + // We need this lock to ensure not to create conn after session exit + e.connLock.Lock() + + if e.closed { + e.connLock.Unlock() + return errors.New("session is closed") + } + + conn, actualAddr, err := e.DialFunc(firstMsg.Addr, firstMsg.Data) + if err != nil { + // Fail fast if DialFunc failed + // (usually indicates the connection has been rejected by the ACL) + e.connLock.Unlock() + // CloseWithErr acquires the connLock again + e.CloseWithErr(err) + return err + } + + e.conn = conn + + if firstMsg.Addr != actualAddr { + // Hook changed the address, enable address override + e.OverrideAddr = actualAddr + e.OriginalAddr = firstMsg.Addr + } + go e.receiveLoop() + + e.connLock.Unlock() + return nil +} + +// receiveLoop receives incoming UDP packets, packs them into UDP messages, +// and sends using the IO. +// Exit when either the underlying UDP connection returns error (e.g. closed), +// or the IO returns error when sending. +func (e *udpSessionEntry) receiveLoop() { + udpBuf := make([]byte, protocol.MaxUDPSize) + msgBuf := make([]byte, protocol.MaxUDPMessageSize) + for { + udpN, rAddr, err := e.conn.ReadFrom(udpBuf) + if err != nil { + e.CloseWithErr(err) + return + } + e.Last.Set(time.Now()) + + if e.OriginalAddr != "" { + // Use the original address in the opposite direction, + // otherwise the QUIC clients or NAT on the client side + // may not treat it as the same UDP session. + rAddr = e.OriginalAddr + } + + msg := &protocol.UDPMessage{ + SessionID: e.ID, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: rAddr, + Data: udpBuf[:udpN], + } + err = sendMessageAutoFrag(e.IO, msgBuf, msg) + if err != nil { + e.CloseWithErr(err) + return + } + } +} + +// sendMessageAutoFrag tries to send a UDP message as a whole first, +// but if it fails due to quic.ErrMessageTooLarge, it tries again by +// fragmenting the message. +func sendMessageAutoFrag(io udpIO, buf []byte, msg *protocol.UDPMessage) error { + err := io.SendMessage(buf, msg) + var errTooLarge *quic.DatagramTooLargeError + if errors.As(err, &errTooLarge) { + // Message too large, try fragmentation + msg.PacketID = uint16(rand.Intn(0xFFFF)) + 1 + fMsgs := frag.FragUDPMessage(msg, int(errTooLarge.MaxDatagramPayloadSize)) + if len(fMsgs) == 0 { + return frag.ErrFragmentationLimit + } + for _, fMsg := range fMsgs { + err := io.SendMessage(buf, &fMsg) + if err != nil { + return err + } + } + return nil + } else { + return err + } +} + +// udpSessionManager manages the lifecycle of UDP sessions. +// Each UDP session is identified by a SessionID, and corresponds to a UDP connection. +// A UDP session is created when a UDP message with a new SessionID is received. +// Similar to standard NAT, a UDP session is destroyed when no UDP message is received +// for a certain period of time (specified by idleTimeout). +type udpSessionManager struct { + io udpIO + eventLogger udpEventLogger + idleTimeout time.Duration + maxSessions int + globalSlots chan struct{} + clientKey string + clientSlots *keyedLimiter + + mutex sync.RWMutex + m map[uint32]*udpSessionEntry +} + +func newUDPSessionManager( + io udpIO, + eventLogger udpEventLogger, + idleTimeout time.Duration, + maxSessions int, + globalSlots chan struct{}, + clientKey string, + clientSlots *keyedLimiter, +) *udpSessionManager { + return &udpSessionManager{ + io: io, + eventLogger: eventLogger, + idleTimeout: idleTimeout, + maxSessions: maxSessions, + globalSlots: globalSlots, + clientKey: clientKey, + clientSlots: clientSlots, + m: make(map[uint32]*udpSessionEntry), + } +} + +// Run runs the session manager main loop. +// Exit and returns error when the underlying io returns error (e.g. closed). +func (m *udpSessionManager) Run() error { + stopCh := make(chan struct{}) + go m.idleCleanupLoop(stopCh) + defer close(stopCh) + defer m.cleanup(false) + + for { + msg, err := m.io.ReceiveMessage() + if err != nil { + return err + } + m.feed(msg) + } +} + +func (m *udpSessionManager) idleCleanupLoop(stopCh <-chan struct{}) { + ticker := time.NewTicker(idleCleanupInterval) + defer ticker.Stop() + for { + select { + case <-ticker.C: + m.cleanup(true) + case <-stopCh: + return + } + } +} + +func (m *udpSessionManager) cleanup(idleOnly bool) { + // We use RLock here as we are only scanning the map, not deleting from it. + m.mutex.RLock() + timeoutEntry := make([]*udpSessionEntry, 0, len(m.m)) + now := time.Now() + for _, entry := range m.m { + if !idleOnly || now.Sub(entry.Last.Get()) > m.idleTimeout { + timeoutEntry = append(timeoutEntry, entry) + } + } + m.mutex.RUnlock() + + for _, entry := range timeoutEntry { + // This eventually calls entry.ExitFunc, + // where the m.mutex will be locked again to remove the entry from the map. + entry.CloseWithErr(nil) + } +} + +func (m *udpSessionManager) feed(msg *protocol.UDPMessage) { + // Reject malformed or excessive fragmentation before allocating session + // state. Legitimate 4096-byte Hysteria datagrams need only a handful of + // fragments at the 1200-byte QUIC datagram size. + if msg.FragCount == 0 || msg.FragID >= msg.FragCount || msg.FragCount > frag.MaxFragments { + return + } + m.mutex.RLock() + entry := m.m[msg.SessionID] + m.mutex.RUnlock() + + // Create a new session if not exists + if entry == nil { + m.mutex.Lock() + entry = m.m[msg.SessionID] + if entry == nil { + if len(m.m) >= m.maxSessions || !tryAcquire(m.globalSlots) { + m.mutex.Unlock() + return + } + if !m.clientSlots.tryAcquire(m.clientKey) { + release(m.globalSlots) + m.mutex.Unlock() + return + } + } + m.mutex.Unlock() + } + + if entry == nil { + dialFunc := func(addr string, firstMsgData []byte) (conn UDPConn, actualAddr string, err error) { + // Call the hook + err = m.io.Hook(firstMsgData, &addr) + if err != nil { + return conn, actualAddr, err + } + actualAddr = addr + // Log the event + m.eventLogger.New(msg.SessionID, addr) + // Dial target + conn, err = m.io.UDP(addr) + return conn, actualAddr, err + } + exitFunc := func(err error) { + // Log the event + m.eventLogger.Close(entry.ID, err) + + // Remove the session from the map + m.mutex.Lock() + delete(m.m, entry.ID) + m.mutex.Unlock() + release(m.globalSlots) + m.clientSlots.release(m.clientKey) + } + + entry = newUDPSessionEntry(msg.SessionID, m.io, dialFunc, exitFunc) + + // Insert the admitted session into the map. feed is called by one Run + // goroutine, while cleanup only removes entries through ExitFunc. + m.mutex.Lock() + m.m[msg.SessionID] = entry + m.mutex.Unlock() + } + + // Feed the message to the session + // Feed (send) errors are ignored for now, + // as some are temporary (e.g. invalid address) + _, _ = entry.Feed(msg) +} + +func tryAcquire(slots chan struct{}) bool { + if slots == nil { + return true + } + select { + case slots <- struct{}{}: + return true + default: + return false + } +} + +func release(slots chan struct{}) { + if slots != nil { + <-slots + } +} + +func (m *udpSessionManager) Count() int { + m.mutex.RLock() + defer m.mutex.RUnlock() + return len(m.m) +} + +// keyedLimiter enforces a shared budget across all QUIC connections with the +// same source key. Counts exist only while at least one resource is live, so +// rotating source addresses cannot grow this map beyond the corresponding +// process-wide resource cap. +type keyedLimiter struct { + mu sync.Mutex + max int + counts map[string]int +} + +func newKeyedLimiter(maximum int) *keyedLimiter { + return &keyedLimiter{max: maximum, counts: make(map[string]int)} +} + +func (l *keyedLimiter) tryAcquire(key string) bool { + if l == nil { + return true + } + l.mu.Lock() + defer l.mu.Unlock() + if l.counts[key] >= l.max { + return false + } + l.counts[key]++ + return true +} + +func (l *keyedLimiter) release(key string) { + if l == nil { + return + } + l.mu.Lock() + defer l.mu.Unlock() + count := l.counts[key] + if count <= 1 { + delete(l.counts, key) + return + } + l.counts[key] = count - 1 +} + +func (l *keyedLimiter) count(key string) int { + if l == nil { + return 0 + } + l.mu.Lock() + defer l.mu.Unlock() + return l.counts[key] +} diff --git a/third_party/hysteria-core/server/udp_test.go b/third_party/hysteria-core/server/udp_test.go new file mode 100644 index 0000000..3edf3bd --- /dev/null +++ b/third_party/hysteria-core/server/udp_test.go @@ -0,0 +1,191 @@ +package server + +import ( + "errors" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.uber.org/goleak" + + "github.com/apernet/hysteria/core/v2/internal/protocol" +) + +func TestUDPSessionManager(t *testing.T) { + io := newMockUDPIO(t) + eventLogger := newMockUDPEventLogger(t) + sm := newUDPSessionManager(io, eventLogger, 2*time.Second, 64, make(chan struct{}, 128), "client", nil) + + msgCh := make(chan *protocol.UDPMessage, 4) + io.EXPECT().ReceiveMessage().RunAndReturn(func() (*protocol.UDPMessage, error) { + m := <-msgCh + if m == nil { + return nil, errors.New("closed") + } + return m, nil + }) + + go sm.Run() + + udpReadFunc := func(addr string, ch chan []byte, b []byte) (int, string, error) { + bs := <-ch + if bs == nil { + return 0, "", errors.New("closed") + } + n := copy(b, bs) + return n, addr, nil + } + + // Test normal session creation & timeout + msg1 := &protocol.UDPMessage{ + SessionID: 1234, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "address1.com:9000", + Data: []byte("hello"), + } + eventLogger.EXPECT().New(msg1.SessionID, msg1.Addr).Return().Once() + udpConn1 := newMockUDPConn(t) + udpConn1Ch := make(chan []byte, 1) + io.EXPECT().Hook(msg1.Data, &msg1.Addr).Return(nil).Once() + io.EXPECT().UDP(msg1.Addr).Return(udpConn1, nil).Once() + udpConn1.EXPECT().WriteTo(msg1.Data, msg1.Addr).Return(5, nil).Once() + udpConn1.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(b []byte) (int, string, error) { + return udpReadFunc(msg1.Addr, udpConn1Ch, b) + }) + io.EXPECT().SendMessage(mock.Anything, &protocol.UDPMessage{ + SessionID: msg1.SessionID, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: msg1.Addr, + Data: []byte("hi back"), + }).Return(nil).Once() + msgCh <- msg1 + udpConn1Ch <- []byte("hi back") + + msg2data := []byte("how are you doing?") + msg2_1 := &protocol.UDPMessage{ + SessionID: 5678, + PacketID: 0, + FragID: 0, + FragCount: 2, + Addr: "address2.net:12450", + Data: msg2data[:6], + } + msg2_2 := &protocol.UDPMessage{ + SessionID: 5678, + PacketID: 0, + FragID: 1, + FragCount: 2, + Addr: "address2.net:12450", + Data: msg2data[6:], + } + + eventLogger.EXPECT().New(msg2_1.SessionID, msg2_1.Addr).Return().Once() + udpConn2 := newMockUDPConn(t) + udpConn2Ch := make(chan []byte, 1) + // On fragmentation, make sure hook gets the whole message + io.EXPECT().Hook(msg2data, &msg2_1.Addr).Return(nil).Once() + io.EXPECT().UDP(msg2_1.Addr).Return(udpConn2, nil).Once() + udpConn2.EXPECT().WriteTo(msg2data, msg2_1.Addr).Return(11, nil).Once() + udpConn2.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(b []byte) (int, string, error) { + return udpReadFunc(msg2_1.Addr, udpConn2Ch, b) + }) + io.EXPECT().SendMessage(mock.Anything, &protocol.UDPMessage{ + SessionID: msg2_1.SessionID, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: msg2_1.Addr, + Data: []byte("im fine"), + }).Return(nil).Once() + msgCh <- msg2_1 + msgCh <- msg2_2 + udpConn2Ch <- []byte("im fine") + + msg3 := &protocol.UDPMessage{ + SessionID: 1234, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "address1.com:9000", + Data: []byte("who are you?"), + } + udpConn1.EXPECT().WriteTo(msg3.Data, msg3.Addr).Return(12, nil).Once() + io.EXPECT().SendMessage(mock.Anything, &protocol.UDPMessage{ + SessionID: msg3.SessionID, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: msg3.Addr, + Data: []byte("im your father"), + }).Return(nil).Once() + msgCh <- msg3 + udpConn1Ch <- []byte("im your father") + + // Make sure timeout works (connections closed & close events emitted) + udpConn1.EXPECT().Close().RunAndReturn(func() error { + close(udpConn1Ch) + return nil + }).Once() + udpConn2.EXPECT().Close().RunAndReturn(func() error { + close(udpConn2Ch) + return nil + }).Once() + eventLogger.EXPECT().Close(msg1.SessionID, nil).Once() + eventLogger.EXPECT().Close(msg2_1.SessionID, nil).Once() + + time.Sleep(3 * time.Second) // Wait for timeout + mock.AssertExpectationsForObjects(t, io, eventLogger, udpConn1, udpConn2) + + // Test UDP connection close error propagation + errUDPClosed := errors.New("UDP connection closed") + msg4 := &protocol.UDPMessage{ + SessionID: 666, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "oh-no.com:27015", + Data: []byte("dont say bye"), + } + eventLogger.EXPECT().New(msg4.SessionID, msg4.Addr).Return().Once() + udpConn4 := newMockUDPConn(t) + io.EXPECT().Hook(msg4.Data, &msg4.Addr).Return(nil).Once() + io.EXPECT().UDP(msg4.Addr).Return(udpConn4, nil).Once() + udpConn4.EXPECT().WriteTo(msg4.Data, msg4.Addr).Return(12, nil).Once() + udpConn4.EXPECT().ReadFrom(mock.Anything).Return(0, "", errUDPClosed).Once() + udpConn4.EXPECT().Close().Return(nil).Once() + eventLogger.EXPECT().Close(msg4.SessionID, errUDPClosed).Once() + msgCh <- msg4 + + time.Sleep(1 * time.Second) + mock.AssertExpectationsForObjects(t, io, eventLogger, udpConn4) + + // Test UDP connection creation error propagation + errUDPIO := errors.New("UDP IO error") + msg5 := &protocol.UDPMessage{ + SessionID: 777, + PacketID: 0, + FragID: 0, + FragCount: 1, + Addr: "callmemaybe.com:15353", + Data: []byte("babe i miss you"), + } + eventLogger.EXPECT().New(msg5.SessionID, msg5.Addr).Return().Once() + io.EXPECT().Hook(msg5.Data, &msg5.Addr).Return(nil).Once() + io.EXPECT().UDP(msg5.Addr).Return(nil, errUDPIO).Once() + eventLogger.EXPECT().Close(msg5.SessionID, errUDPIO).Once() + msgCh <- msg5 + + time.Sleep(1 * time.Second) + mock.AssertExpectationsForObjects(t, io, eventLogger) + + // Leak checks + close(msgCh) // This will return error from ReceiveMessage(), should stop the session manager + time.Sleep(1 * time.Second) // Wait one more second just to be sure + assert.Zero(t, sm.Count(), "session count should be 0") + goleak.VerifyNone(t) +} diff --git a/third_party/quic-go/.clusterfuzzlite/Dockerfile b/third_party/quic-go/.clusterfuzzlite/Dockerfile new file mode 100644 index 0000000..9c00ced --- /dev/null +++ b/third_party/quic-go/.clusterfuzzlite/Dockerfile @@ -0,0 +1,5 @@ +FROM gcr.io/oss-fuzz-base/base-builder-go:v1 + +COPY . $SRC/quic-go +WORKDIR $SRC/quic-go +COPY .clusterfuzzlite/build.sh $SRC/ diff --git a/third_party/quic-go/.clusterfuzzlite/build.sh b/third_party/quic-go/.clusterfuzzlite/build.sh new file mode 100644 index 0000000..d11d10b --- /dev/null +++ b/third_party/quic-go/.clusterfuzzlite/build.sh @@ -0,0 +1,38 @@ +#!/bin/bash + +set -euo pipefail + +go version +go env + +build_native_go_fuzzer() { + local pkg=$1 + local fuzz=$2 + local name=$3 + local corpus_dir="${WORK:-/tmp}/quic-go-seed-corpus/$name" + local corpus_zip="$OUT/${name}_seed_corpus.zip" + + # FUZZ_CORPUS_DIR makes go-ossfuzz-seeds write each f.Add seed as a raw + # libFuzzer corpus file. OSS-Fuzz picks up _seed_corpus.zip from + # $OUT and unpacks it next to the fuzzer binary. + rm -rf "$corpus_dir" + mkdir -p "$corpus_dir" + FUZZ_CORPUS_DIR="$corpus_dir" go test "$pkg" -run "^${fuzz}$" -count=1 -v + + rm -f "$corpus_zip" + corpus_files=$(find "$corpus_dir" -type f | wc -l) + echo "$name: generated $corpus_files corpus files" + if [[ "$corpus_files" -gt 0 ]]; then + (cd "$corpus_dir" && zip -q -r "$corpus_zip" .) + fi + + compile_native_go_fuzzer_v2 "$pkg" "$fuzz" "$name" +} + +build_native_go_fuzzer github.com/quic-go/quic-go/internal/wire FuzzFrames frame_fuzzer_v2 +build_native_go_fuzzer github.com/quic-go/quic-go/internal/wire FuzzTransportParameters transportparameter_fuzzer_v2 +build_native_go_fuzzer github.com/quic-go/quic-go/http3 FuzzFrameParser http3_frame_fuzzer +build_native_go_fuzzer github.com/quic-go/quic-go/internal/wire FuzzHeaderParser header_fuzzer_v2 +build_native_go_fuzzer github.com/quic-go/quic-go/internal/handshake FuzzHandshake handshake_fuzzer_v2 +build_native_go_fuzzer github.com/quic-go/quic-go FuzzFrameSorter frame_sorter_fuzzer +build_native_go_fuzzer github.com/quic-go/quic-go/http3 FuzzHeaderParsing http3_header_parsing_fuzzer diff --git a/third_party/quic-go/.clusterfuzzlite/project.yaml b/third_party/quic-go/.clusterfuzzlite/project.yaml new file mode 100644 index 0000000..4f2ee4d --- /dev/null +++ b/third_party/quic-go/.clusterfuzzlite/project.yaml @@ -0,0 +1 @@ +language: go diff --git a/third_party/quic-go/.githooks/README.md b/third_party/quic-go/.githooks/README.md new file mode 100644 index 0000000..e38700c --- /dev/null +++ b/third_party/quic-go/.githooks/README.md @@ -0,0 +1,8 @@ +# Git Hooks + +This directory contains useful Git hooks for working with quic-go. + +Install them by running +```bash +git config core.hooksPath .githooks +``` diff --git a/third_party/quic-go/.githooks/pre-commit b/third_party/quic-go/.githooks/pre-commit new file mode 100644 index 0000000..0e3c572 --- /dev/null +++ b/third_party/quic-go/.githooks/pre-commit @@ -0,0 +1,34 @@ +#!/bin/bash + +# Check that test files don't contain focussed test cases. +errored=false +for f in $(git diff --diff-filter=d --cached --name-only); do + if [[ $f != *_test.go ]]; then continue; fi + output=$(git show :"$f" | grep -n -e "FIt(" -e "FContext(" -e "FDescribe(") + if [ $? -eq 0 ]; then + echo "$f contains a focussed test:" + echo "$output" + echo "" + errored=true + fi +done + +pushd ./integrationtests/gomodvendor > /dev/null +go mod tidy +if [[ -n $(git diff --diff-filter=d --name-only -- "go.mod" "go.sum") ]]; then + echo "go.mod / go.sum in integrationtests/gomodvendor not tidied" + errored=true +fi +popd > /dev/null + +# Check that all Go files are properly gofumpt-ed. +output=$(gofumpt -d $(git diff --diff-filter=d --cached --name-only -- '*.go')) +if [ -n "$output" ]; then + echo "Found files that are not properly gofumpt-ed." + echo "$output" + errored=true +fi + +if [ "$errored" = true ]; then + exit 1 +fi diff --git a/third_party/quic-go/.github/FUNDING.yml b/third_party/quic-go/.github/FUNDING.yml new file mode 100644 index 0000000..7de30a1 --- /dev/null +++ b/third_party/quic-go/.github/FUNDING.yml @@ -0,0 +1,13 @@ +# These are supported funding model platforms + +github: [marten-seemann] # Replace with up to 4 GitHub Sponsors-enabled usernames e.g., [user1, user2] +patreon: # Replace with a single Patreon username +open_collective: # Replace with a single Open Collective username +ko_fi: # Replace with a single Ko-fi username +tidelift: # Replace with a single Tidelift platform-name/package-name e.g., npm/babel +community_bridge: # Replace with a single Community Bridge project-name e.g., cloud-foundry +liberapay: # Replace with a single Liberapay username +issuehunt: # Replace with a single IssueHunt username +otechie: # Replace with a single Otechie username +lfx_crowdfunding: # Replace with a single LFX Crowdfunding project-name e.g., cloud-foundry +custom: # Replace with up to 4 custom sponsorship URLs e.g., ['link1', 'link2'] diff --git a/third_party/quic-go/.github/dependabot.yml b/third_party/quic-go/.github/dependabot.yml new file mode 100644 index 0000000..5ace460 --- /dev/null +++ b/third_party/quic-go/.github/dependabot.yml @@ -0,0 +1,6 @@ +version: 2 +updates: + - package-ecosystem: "github-actions" + directory: "/" + schedule: + interval: "weekly" diff --git a/third_party/quic-go/.github/workflows/build-interop-docker.yml b/third_party/quic-go/.github/workflows/build-interop-docker.yml new file mode 100644 index 0000000..2bdcf0e --- /dev/null +++ b/third_party/quic-go/.github/workflows/build-interop-docker.yml @@ -0,0 +1,51 @@ +name: Build interop Docker image + +permissions: read-all + +on: + push: + branches: + - master + tags: + - 'v*' + pull_request: + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: ${{ github.event_name == 'push' && github.ref != 'refs/heads/master' }} + +jobs: + interop: + runs-on: ${{ fromJSON(vars['DOCKER_RUNNER_UBUNTU'] || '"ubuntu-latest"') }} + timeout-minutes: 30 + steps: + - uses: actions/checkout@v7 + - name: Set up QEMU + uses: docker/setup-qemu-action@96fe6ef7f33517b61c61be40b68a1882f3264fb8 # v4.2.0 + - name: Set up Docker Buildx + uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 + with: + platforms: linux/amd64,linux/arm64 + - name: Login to Docker Hub + if: github.event_name == 'push' + uses: docker/login-action@af1e73f918a031802d376d3c8bbc3fe56130a9b0 # v4.4.0 + with: + username: ${{ secrets.DOCKER_USERNAME }} + password: ${{ secrets.DOCKER_PASSWORD }} + - name: set tag name + id: tag + # Tagged releases won't be picked up by the interop runner automatically, + # but they can be useful when debugging regressions. + run: | + if [[ $GITHUB_REF == refs/tags/* ]]; then + echo "tag=${GITHUB_REF#refs/tags/}" | tee -a $GITHUB_OUTPUT; + else + echo 'tag=latest' | tee -a $GITHUB_OUTPUT; + fi + - uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0 + with: + context: "." + file: "interop/Dockerfile" + platforms: linux/amd64,linux/arm64 + push: ${{ github.event_name == 'push' }} + tags: martenseemann/quic-go-interop:${{ steps.tag.outputs.tag }} diff --git a/third_party/quic-go/.github/workflows/clusterfuzz-coverage.yml b/third_party/quic-go/.github/workflows/clusterfuzz-coverage.yml new file mode 100644 index 0000000..5b6d342 --- /dev/null +++ b/third_party/quic-go/.github/workflows/clusterfuzz-coverage.yml @@ -0,0 +1,106 @@ +name: ClusterFuzz coverage +on: + schedule: + - cron: '12 3,11,19 * * *' + workflow_dispatch: + +permissions: + contents: read + +jobs: + coverage: + runs-on: ubuntu-latest + env: + DOCKER_DEFAULT_PLATFORM: linux/amd64 + CLUSTERFUZZ_CORPUS_BUCKET: quic-go-corpus.clusterfuzz-external.appspot.com + steps: + - uses: actions/checkout@v7 + - uses: actions/checkout@v7 + with: + repository: google/oss-fuzz + path: oss-fuzz + - uses: google-github-actions/setup-gcloud@aa5489c8933f4cc7a4f7d45035b3b1440c9c10db # v3.0.1 + - name: Configure Google Cloud credentials + run: | + cat > "$RUNNER_TEMP/gcp-ossfuzz-user-credentials.json" <<'EOF' + ${{ secrets.GCP_OSSFUZZ_USER_CREDENTIALS }} + EOF + echo "CLOUDSDK_AUTH_CREDENTIAL_FILE_OVERRIDE=$RUNNER_TEMP/gcp-ossfuzz-user-credentials.json" >> "$GITHUB_ENV" + - name: Download ClusterFuzz corpus + run: | + set -euo pipefail + + targets=( + frame_fuzzer_v2 + transportparameter_fuzzer_v2 + http3_frame_fuzzer + header_fuzzer_v2 + handshake_fuzzer_v2 + frame_sorter_fuzzer + http3_header_parsing_fuzzer + ) + + mkdir -p oss-fuzz/build/corpus/quic-go + for target in "${targets[@]}"; do + source_dir="gs://${CLUSTERFUZZ_CORPUS_BUCKET}/libFuzzer/quic-go_${target}" + corpus_dir="oss-fuzz/build/corpus/quic-go/${target}" + + echo "$target: downloading live corpus" + rm -rf "$corpus_dir" + mkdir -p "$corpus_dir" + + gcloud storage cp --recursive "$source_dir/*" "$corpus_dir" || echo "$target: no live corpus found" + done + + corpus_files=$(find oss-fuzz/build/corpus/quic-go -type f | wc -l | tr -d ' ') + if [ "$corpus_files" = 0 ]; then + echo "no ClusterFuzz corpus files downloaded" + exit 1 + fi + - name: Summarize ClusterFuzz corpus + run: | + set -euo pipefail + + corpus_root=oss-fuzz/build/corpus/quic-go + { + echo "## ClusterFuzz corpus" + echo + echo "| Fuzzer | Corpus files downloaded |" + echo "| --- | ---: |" + + for corpus_dir in "$corpus_root"/*; do + [ -d "$corpus_dir" ] || continue + + target=$(basename "$corpus_dir") + corpus_files=$(find "$corpus_dir" -type f | wc -l | tr -d ' ') + if [ "$corpus_files" = 0 ]; then + corpus_files="---" + fi + echo "| \`$target\` | $corpus_files |" + done + } >> "$GITHUB_STEP_SUMMARY" + - name: Build coverage fuzzers + working-directory: oss-fuzz + run: | + python3 infra/helper.py build_image --no-pull quic-go + python3 infra/helper.py build_fuzzers --sanitizer coverage --mount_path /src/quic-go quic-go "$GITHUB_WORKSPACE" + - name: Generate coverage + working-directory: oss-fuzz + run: | + rm -f build/out/quic-go/qpack_decode_fuzzer + python3 infra/helper.py coverage --no-corpus-download --no-serve quic-go + - name: Prepare Codecov coverage + working-directory: oss-fuzz + run: | + sed "s#^/out/src/quic-go/#${GITHUB_WORKSPACE}/#" \ + build/out/quic-go/fuzz.cov > "$GITHUB_WORKSPACE/clusterfuzz.coverprofile" + test -s "$GITHUB_WORKSPACE/clusterfuzz.coverprofile" + - name: Upload coverage to Codecov + if: ${{ !cancelled() }} + uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0 + with: + disable_search: true + files: clusterfuzz.coverprofile + flags: clusterfuzz + name: ClusterFuzz fuzzing + token: ${{ secrets.CODECOV_TOKEN }} diff --git a/third_party/quic-go/.github/workflows/clusterfuzz-lite-batch.yml b/third_party/quic-go/.github/workflows/clusterfuzz-lite-batch.yml new file mode 100644 index 0000000..ea6b30b --- /dev/null +++ b/third_party/quic-go/.github/workflows/clusterfuzz-lite-batch.yml @@ -0,0 +1,87 @@ +name: ClusterFuzzLite batch fuzzing +on: + schedule: + - cron: '0 0/6 * * *' + +permissions: + contents: read + security-events: write + +jobs: + batch-fuzzing: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + sanitizer: + - address + steps: + - name: Build Fuzzers (${{ matrix.sanitizer }}) + id: build + uses: google/clusterfuzzlite/actions/build_fuzzers@884713a6c30a92e5e8544c39945cd7cb630abcd1 # v1 + with: + language: go + sanitizer: ${{ matrix.sanitizer }} + - name: Run Fuzzers (${{ matrix.sanitizer }}) + id: run + uses: google/clusterfuzzlite/actions/run_fuzzers@884713a6c30a92e5e8544c39945cd7cb630abcd1 # v1 + with: + github-token: ${{ secrets.GITHUB_TOKEN }} + fuzz-seconds: 3600 + mode: 'batch' + sanitizer: ${{ matrix.sanitizer }} + output-sarif: true + storage-repo: https://${{ secrets.CLUSTERFUZZ_LITE_STORAGE }}@github.com/quic-go/clusterfuzzlite-storage.git + storage-repo-branch: master + storage-repo-branch-coverage: gh-pages + + coverage: + needs: batch-fuzzing + runs-on: ubuntu-latest + env: + DOCKER_DEFAULT_PLATFORM: linux/amd64 + steps: + - uses: actions/checkout@v7 + - uses: actions/checkout@v7 + with: + repository: google/oss-fuzz + path: oss-fuzz + - uses: actions/checkout@v7 + with: + repository: quic-go/clusterfuzzlite-storage + ref: master + path: clusterfuzzlite-storage + token: ${{ secrets.CLUSTERFUZZ_LITE_STORAGE }} + - name: List fuzz targets + run: find clusterfuzzlite-storage/corpus -mindepth 1 -maxdepth 1 -type d -printf '%f\n' | sort + - name: Prepare corpus + run: | + mkdir -p oss-fuzz/build/corpus + mv clusterfuzzlite-storage/corpus oss-fuzz/build/corpus/quic-go + - name: Use ClusterFuzzLite build script + run: | + cp .clusterfuzzlite/build.sh oss-fuzz/projects/quic-go/build.sh + sed -i '4i cd "$SRC/quic-go"' oss-fuzz/projects/quic-go/build.sh + - name: Build coverage fuzzers + working-directory: oss-fuzz + run: | + python3 infra/helper.py build_image --no-pull quic-go + python3 infra/helper.py build_fuzzers --sanitizer coverage --mount_path /src/quic-go quic-go "$GITHUB_WORKSPACE" + - name: Generate coverage + working-directory: oss-fuzz + run: python3 infra/helper.py coverage --no-corpus-download --no-serve quic-go + - name: Prepare Codecov coverage + working-directory: oss-fuzz + run: | + sed "s#^/out/src/quic-go/#${GITHUB_WORKSPACE}/#" \ + build/out/quic-go/fuzz.cov > "$GITHUB_WORKSPACE/clusterfuzz-lite-batch.coverprofile" + test -s "$GITHUB_WORKSPACE/clusterfuzz-lite-batch.coverprofile" + - name: Upload coverage to Codecov + if: ${{ !cancelled() }} + uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0 + with: + disable_search: true + files: clusterfuzz-lite-batch.coverprofile + flags: clusterfuzz-lite-batch + name: ClusterFuzzLite batch fuzzing + token: ${{ secrets.CODECOV_TOKEN }} diff --git a/third_party/quic-go/.github/workflows/clusterfuzz-lite-pr.yml b/third_party/quic-go/.github/workflows/clusterfuzz-lite-pr.yml new file mode 100644 index 0000000..afdb417 --- /dev/null +++ b/third_party/quic-go/.github/workflows/clusterfuzz-lite-pr.yml @@ -0,0 +1,51 @@ +name: ClusterFuzzLite PR fuzzing +on: + pull_request: + paths: + - '**' + +permissions: + contents: read + security-events: write + +jobs: + PR: + if: ${{ github.event.pull_request.user.login != 'dependabot[bot]' }} + env: + # Forked PRs don't receive repository secrets, so they run without corpus storage. + CLUSTERFUZZ_LITE_STORAGE_REPO: >- + ${{ + github.event.pull_request.head.repo.id == github.event.pull_request.base.repo.id && + format('https://{0}@github.com/quic-go/clusterfuzzlite-storage.git', secrets.CLUSTERFUZZ_LITE_STORAGE) || '' + }} + runs-on: ${{ fromJSON(vars['CLUSTERFUZZ_LITE_RUNNER_UBUNTU'] || '"ubuntu-latest"') }} + concurrency: + group: ${{ github.workflow }}-${{ matrix.sanitizer }}-${{ github.ref }} + cancel-in-progress: true + strategy: + fail-fast: false + matrix: + sanitizer: + - address + steps: + - name: Build Fuzzers (${{ matrix.sanitizer }}) + uses: google/clusterfuzzlite/actions/build_fuzzers@884713a6c30a92e5e8544c39945cd7cb630abcd1 # v1 + with: + language: go + github-token: ${{ secrets.GITHUB_TOKEN }} + sanitizer: ${{ matrix.sanitizer }} + storage-repo: ${{ env.CLUSTERFUZZ_LITE_STORAGE_REPO }} + storage-repo-branch: master + storage-repo-branch-coverage: gh-pages + - name: Run Fuzzers (${{ matrix.sanitizer }}) + uses: google/clusterfuzzlite/actions/run_fuzzers@884713a6c30a92e5e8544c39945cd7cb630abcd1 # v1 + with: + github-token: ${{ secrets.GITHUB_TOKEN }} + fuzz-seconds: 480 + mode: 'code-change' + sanitizer: ${{ matrix.sanitizer }} + output-sarif: true + parallel-fuzzing: true + storage-repo: ${{ env.CLUSTERFUZZ_LITE_STORAGE_REPO }} + storage-repo-branch: master + storage-repo-branch-coverage: gh-pages diff --git a/third_party/quic-go/.github/workflows/clusterfuzz-lite-prune.yml b/third_party/quic-go/.github/workflows/clusterfuzz-lite-prune.yml new file mode 100644 index 0000000..cf420fd --- /dev/null +++ b/third_party/quic-go/.github/workflows/clusterfuzz-lite-prune.yml @@ -0,0 +1,29 @@ +name: ClusterFuzzLite pruning +on: + schedule: + - cron: '0 2 * * *' + +permissions: + contents: read + security-events: write + +jobs: + Pruning: + runs-on: ubuntu-latest + steps: + - name: Build Fuzzers + id: build + uses: google/clusterfuzzlite/actions/build_fuzzers@884713a6c30a92e5e8544c39945cd7cb630abcd1 # v1 + with: + language: go + - name: Run Fuzzers + id: run + uses: google/clusterfuzzlite/actions/run_fuzzers@884713a6c30a92e5e8544c39945cd7cb630abcd1 # v1 + with: + github-token: ${{ secrets.GITHUB_TOKEN }} + fuzz-seconds: 600 + mode: 'prune' + output-sarif: true + storage-repo: https://${{ secrets.CLUSTERFUZZ_LITE_STORAGE }}@github.com/quic-go/clusterfuzzlite-storage.git + storage-repo-branch: master + storage-repo-branch-coverage: gh-pages diff --git a/third_party/quic-go/.github/workflows/codspeed.yml b/third_party/quic-go/.github/workflows/codspeed.yml new file mode 100644 index 0000000..f30db5e --- /dev/null +++ b/third_party/quic-go/.github/workflows/codspeed.yml @@ -0,0 +1,30 @@ +name: Benchmarks + +permissions: + contents: read + id-token: write + +on: + push: + branches: [master] + pull_request: + schedule: + - cron: "0 0 * * *" + workflow_dispatch: + +jobs: + benchmarks: + name: Benchmarks + runs-on: codspeed-macro + timeout-minutes: 180 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: "1.26.x" + check-latest: true + - name: Run the benchmarks + uses: CodSpeedHQ/action@f99becdce5e5d51fd556489ebef684f4ecfd6286 # v4.18.5 + with: + mode: walltime + run: go test -run=^$ -bench=. ./integrationtests/self diff --git a/third_party/quic-go/.github/workflows/cross-compile.sh b/third_party/quic-go/.github/workflows/cross-compile.sh new file mode 100644 index 0000000..cd52622 --- /dev/null +++ b/third_party/quic-go/.github/workflows/cross-compile.sh @@ -0,0 +1,33 @@ +#!/bin/bash + +set -e + +dist="$1" +goos=$(echo "$dist" | cut -d "/" -f1) +goarch=$(echo "$dist" | cut -d "/" -f2) + +# cross-compiling for android is a pain... +if [[ "$goos" == "android" ]]; then exit; fi +# iOS builds require Cgo, see https://github.com/golang/go/issues/43343 +# Cgo would then need a C cross compilation setup. Not worth the hassle. +if [[ "$goos" == "ios" ]]; then exit; fi + +# Write all log output to a temporary file instead of to stdout. +# That allows running this script in parallel, while preserving the correct order of the output. +log_file=$(mktemp) + +error_handler() { + cat "$log_file" >&2 + rm "$log_file" + exit 1 +} + +trap 'error_handler' ERR + +echo "$dist" >> "$log_file" +out="main-$goos-$goarch" +GOOS=$goos GOARCH=$goarch go build -o $out example/main.go >> "$log_file" 2>&1 +rm $out + +cat "$log_file" +rm "$log_file" diff --git a/third_party/quic-go/.github/workflows/cross-compile.yml b/third_party/quic-go/.github/workflows/cross-compile.yml new file mode 100644 index 0000000..0861b9c --- /dev/null +++ b/third_party/quic-go/.github/workflows/cross-compile.yml @@ -0,0 +1,50 @@ +on: [push, pull_request] + +permissions: read-all + +jobs: + crosscompile: + permissions: + actions: write + contents: read + strategy: + fail-fast: false + matrix: + go: [ "1.25.x", "1.26.x", "1.27.0-rc.1" ] + runs-on: ${{ fromJSON(vars['CROSS_COMPILE_RUNNER_UBUNTU'] || '"ubuntu-latest"') }} + name: "Cross Compilation (Go ${{matrix.go}})" + timeout-minutes: 30 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: ${{ matrix.go }} + check-latest: true + - name: Get Date + id: get-date + run: echo "date=$(/bin/date -u "+%Y%m%d")" >> $GITHUB_OUTPUT + - name: Load Go build cache + id: load-go-cache + uses: actions/cache/restore@v6 + with: + path: ~/.cache/go-build + key: go-${{ matrix.go }}-crosscompile-${{ steps.get-date.outputs.date }} + restore-keys: go-${{ matrix.go }}-crosscompile- + - name: Install build utils + run: | + sudo apt-get update + sudo apt-get install -y gcc-multilib + - name: Install dependencies + run: go build example/main.go + - name: Run cross compilation + # run in parallel on as many cores as are available on the machine + run: go tool dist list | xargs -I % -P "$(nproc)" .github/workflows/cross-compile.sh % + - name: Save Go build cache + # only store cache when on master + if: github.event_name == 'push' && github.ref_name == 'master' + uses: actions/cache/save@v6 + with: + path: ~/.cache/go-build + # Caches are immutable, so we only update it once per day (at most). + # See https://github.com/actions/cache/blob/main/tips-and-workarounds.md#update-a-cache + key: go-${{ matrix.go }}-crosscompile-${{ steps.get-date.outputs.date }} diff --git a/third_party/quic-go/.github/workflows/go-generate.sh b/third_party/quic-go/.github/workflows/go-generate.sh new file mode 100644 index 0000000..fab97c7 --- /dev/null +++ b/third_party/quic-go/.github/workflows/go-generate.sh @@ -0,0 +1,18 @@ +#!/usr/bin/env bash + +set -e + +# delete all go-generated files (that adhere to the comment convention) +git ls-files -z | grep --include \*.go -lrIZ "^// Code generated .* DO NOT EDIT\.$" | tr '\0' '\n' | xargs rm -f + +# First regenerate sys_conn_buffers_write.go. +# If it doesn't exist, the following mockgen calls will fail. +go generate -run "sys_conn_buffers_write.go" +# now generate everything +go generate ./... + +# Check if any files were changed +git diff --exit-code || ( + echo "Generated files are not up to date. Please run 'go generate ./...' and commit the changes." + exit 1 +) diff --git a/third_party/quic-go/.github/workflows/govulncheck.yml b/third_party/quic-go/.github/workflows/govulncheck.yml new file mode 100644 index 0000000..8b69155 --- /dev/null +++ b/third_party/quic-go/.github/workflows/govulncheck.yml @@ -0,0 +1,19 @@ +on: + schedule: + - cron: "17 4 * * *" + workflow_dispatch: + +permissions: read-all + +jobs: + govulncheck: + runs-on: ubuntu-latest + timeout-minutes: 15 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: "1.26.x" + check-latest: true + - name: Run govulncheck + run: go run golang.org/x/vuln/cmd/govulncheck@latest ./... diff --git a/third_party/quic-go/.github/workflows/integration.yml b/third_party/quic-go/.github/workflows/integration.yml new file mode 100644 index 0000000..8a6f7cd --- /dev/null +++ b/third_party/quic-go/.github/workflows/integration.yml @@ -0,0 +1,89 @@ +on: [push, pull_request] + +permissions: read-all + +jobs: + integration: + strategy: + fail-fast: false + matrix: + os: [ "ubuntu" ] + go: [ "1.25.x", "1.26.x", "1.27.0-rc.1" ] + race: [ false ] + include: + - os: "ubuntu" + go: "1.25.x" + race: true + - os: "windows" + go: "1.25.x" + race: false + - os: "macos" + go: "1.25.x" + race: false + runs-on: ${{ fromJSON(vars[format('INTEGRATION_RUNNER_{0}', matrix.os)] || format('"{0}-latest"', matrix.os)) }} + timeout-minutes: 30 + defaults: + run: + shell: bash # by default Windows uses PowerShell, which uses a different syntax for setting environment variables + env: + DEBUG: false # set this to true to export qlogs and save them as artifacts + TIMESCALE_FACTOR: 3 + name: "Integration (${{ matrix.os }}, Go ${{ matrix.go }}${{ matrix.race && ', race' || '' }})" + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: ${{ matrix.go }} + check-latest: true + - name: Install go-junit-report + run: go install github.com/jstemmer/go-junit-report/v2@v2.1.0 + - name: Set qlogger + if: env.DEBUG == 'true' + run: echo "QLOGFLAG= -qlog" >> $GITHUB_ENV + - name: Enable race detector + if: ${{ matrix.race }} + run: echo "RACEFLAG= -race" >> $GITHUB_ENV + - run: go version + - name: Run tools tests + run: go test ${{ env.RACEFLAG }} -v -timeout 30s -shuffle=on ./integrationtests/tools/... 2>&1 | go-junit-report -set-exit-code -iocopy -out report_tools.xml + - name: Run version negotiation tests + run: go test ${{ env.RACEFLAG }} -v -timeout 30s -shuffle=on ./integrationtests/versionnegotiation ${{ env.QLOGFLAG }} 2>&1 | go-junit-report -set-exit-code -iocopy -out report_versionnegotiation.xml + - name: Run FIPS 140 tests + if: ${{ matrix.go != '1.25.x' && (success() || failure()) }} + working-directory: integrationtests/fips + run: go test -v -timeout 1m -shuffle=on . 2>&1 | go-junit-report -set-exit-code -iocopy -out ../../report_fips.xml + - name: Run self tests, using QUIC v1 + if: success() || failure() # run this step even if the previous one failed + run: go test ${{ env.RACEFLAG }} -v -timeout 5m -shuffle=on ./integrationtests/self -version=1 ${{ env.QLOGFLAG }} 2>&1 | go-junit-report -set-exit-code -iocopy -out report_self.xml + - name: Run self tests, using QUIC v2 + if: ${{ !matrix.race && (success() || failure()) }} # run this step even if the previous one failed + run: go test ${{ env.RACEFLAG }} -v -timeout 5m -shuffle=on ./integrationtests/self -version=2 ${{ env.QLOGFLAG }} 2>&1 | go-junit-report -set-exit-code -iocopy -out report_self_v2.xml + - name: Run self tests, with GSO disabled + if: ${{ matrix.os == 'ubuntu' && (success() || failure()) }} # run this step even if the previous one failed + env: + QUIC_GO_DISABLE_GSO: true + run: go test ${{ env.RACEFLAG }} -v -timeout 5m -shuffle=on ./integrationtests/self -version=1 ${{ env.QLOGFLAG }} 2>&1 | go-junit-report -set-exit-code -iocopy -out report_self_nogso.xml + - name: Run self tests, with ECN disabled + if: ${{ !matrix.race && matrix.os == 'ubuntu' && (success() || failure()) }} # run this step even if the previous one failed + env: + QUIC_GO_DISABLE_ECN: true + run: go test ${{ env.RACEFLAG }} -v -timeout 5m -shuffle=on ./integrationtests/self -version=1 ${{ env.QLOGFLAG }} 2>&1 | go-junit-report -set-exit-code -iocopy -out report_self_noecn.xml + - name: Run benchmarks + if: ${{ !matrix.race }} + run: go test -v -run=^$ -timeout 5m -shuffle=on -bench=. ./integrationtests/self + - name: save qlogs + if: ${{ always() && env.DEBUG == 'true' }} + uses: actions/upload-artifact@v7 + with: + name: qlogs-${{ matrix.os }}-go${{ matrix.go }}-race${{ matrix.race }} + path: integrationtests/self/*.qlog + retention-days: 7 + - name: Upload report to Codecov + if: ${{ !cancelled() && !matrix.race }} + uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0 + with: + report_type: test_results + name: Unit tests + files: report_tools.xml,report_versionnegotiation.xml,report_self.xml,report_self_v2.xml,report_self_nogso.xml,report_self_noecn.xml,report_fips.xml + env_vars: OS,GO + token: ${{ secrets.CODECOV_TOKEN }} diff --git a/third_party/quic-go/.github/workflows/lint.yml b/third_party/quic-go/.github/workflows/lint.yml new file mode 100644 index 0000000..a0c00b4 --- /dev/null +++ b/third_party/quic-go/.github/workflows/lint.yml @@ -0,0 +1,101 @@ +on: [push, pull_request] + +permissions: read-all + +jobs: + check: + runs-on: ubuntu-latest + timeout-minutes: 15 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: "1.26.x" + check-latest: true + - name: Check for //go:build ignore in .go files + run: | + IGNORED_FILES=$(grep -rl '//go:build ignore' . --include='*.go') || true + if [ -n "$IGNORED_FILES" ]; then + echo "::error::Found ignored Go files: $IGNORED_FILES" + exit 1 + fi + - name: Check that go.mod is tidied + if: success() || failure() # run this step even if the previous one failed + run: go mod tidy -diff + - name: Check that FIPS go.mod is tidied + if: success() || failure() # run this step even if the previous one failed + working-directory: integrationtests/fips + run: go mod tidy -diff + - name: Run go fix + if: success() || failure() # run this step even if the previous one failed + run: go fix -diff ./... + - name: Run code generators + if: success() || failure() # run this step even if the previous one failed + run: .github/workflows/go-generate.sh + - name: Check that go mod vendor works + if: success() || failure() # run this step even if the previous one failed + run: | + cd integrationtests/gomodvendor + go mod vendor + - name: Run gcassert + if: success() || failure() # run this step even if the previous one failed + run: go tool gcassert ./... + # Only run govulncheck on pull requests, not on pushes (including merge commits to master). + # govulncheck queries the vulnerability database at runtime, so a newly disclosed vulnerability + # could otherwise cause the post-merge run to fail even though nothing in the repository changed, + # breaking master through no fault of the merged PR. + - name: Run govulncheck + if: ${{ (success() || failure()) && github.event_name == 'pull_request' }} + run: go run golang.org/x/vuln/cmd/govulncheck@latest ./... + golangci-lint: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + go: [ "1.25.x", "1.26.x" ] + env: + GOLANGCI_LINT_VERSION: v2.11.4 + name: golangci-lint (Go ${{ matrix.go }}) + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: ${{ matrix.go }} + check-latest: true + - name: golangci-lint (Linux) + uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v9.3.0 + with: + args: --timeout=3m + version: ${{ env.GOLANGCI_LINT_VERSION }} + - name: golangci-lint (Windows) + if: success() || failure() # run this step even if the previous one failed + uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v9.3.0 + env: + GOOS: "windows" + with: + args: --timeout=3m + version: ${{ env.GOLANGCI_LINT_VERSION }} + - name: golangci-lint (OSX) + if: success() || failure() # run this step even if the previous one failed + uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v9.3.0 + env: + GOOS: "darwin" + with: + args: --timeout=3m + version: ${{ env.GOLANGCI_LINT_VERSION }} + - name: golangci-lint (FreeBSD) + if: success() || failure() # run this step even if the previous one failed + uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v9.3.0 + env: + GOOS: "freebsd" + with: + args: --timeout=3m + version: ${{ env.GOLANGCI_LINT_VERSION }} + - name: golangci-lint (others) + if: success() || failure() # run this step even if the previous one failed + uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v9.3.0 + env: + GOOS: "solaris" # some OS that we don't have any build tags for + with: + args: --timeout=3m + version: ${{ env.GOLANGCI_LINT_VERSION }} diff --git a/third_party/quic-go/.github/workflows/unit.yml b/third_party/quic-go/.github/workflows/unit.yml new file mode 100644 index 0000000..747240f --- /dev/null +++ b/third_party/quic-go/.github/workflows/unit.yml @@ -0,0 +1,75 @@ +on: [push, pull_request] + +permissions: read-all + +jobs: + unit: + strategy: + fail-fast: false + matrix: + os: [ "ubuntu", "windows", "macos" ] + go: [ "1.25.x", "1.26.x", "1.27.0-rc.1" ] + runs-on: ${{ fromJSON(vars[format('UNIT_RUNNER_{0}', matrix.os)] || format('"{0}-latest"', matrix.os)) }} + name: Unit tests (${{ matrix.os}}, Go ${{ matrix.go }}) + timeout-minutes: 30 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: ${{ matrix.go }} + check-latest: true + - run: go version + - name: Install go-junit-report + run: go install github.com/jstemmer/go-junit-report/v2@v2.1.0 + - name: Remove integrationtests + shell: bash + run: git rm -r --cached integrationtests && rm -rf integrationtests + - name: Run tests + env: + TIMESCALE_FACTOR: 10 + run: go test -v -shuffle on -cover -coverprofile coverage.txt ./... 2>&1 | go-junit-report -set-exit-code -iocopy -out report.xml + - name: Run tests as root + if: ${{ matrix.os == 'ubuntu' }} + env: + TIMESCALE_FACTOR: 10 + FILE: sys_conn_helper_linux_test.go + run: | + test -f $FILE # make sure the file actually exists + TEST_NAMES=$(grep '^func Test' "$FILE" | sed 's/^func \([A-Za-z0-9_]*\)(.*/\1/' | tr '\n' '|') + go test -c -cover -tags root -o quic-go.test . + sudo ./quic-go.test -test.v -test.run "${TEST_NAMES%|}" -test.coverprofile coverage-root.txt 2>&1 | go-junit-report -set-exit-code -iocopy -package-name github.com/quic-go/quic-go -out report_root.xml + rm quic-go.test + - name: Run tests with race detector + if: ${{ matrix.os == 'ubuntu' }} # speed things up. Windows and OSX VMs are slow + env: + TIMESCALE_FACTOR: 20 + run: go test -v -shuffle on ./... + - name: Run handshake tests in FIPS140 mode + if: ${{ matrix.go != '1.25.x' }} + env: + GODEBUG: fips140=only + run: go test -v ./internal/handshake -run 'TestToken|TestRetry|TestInitial|TestDecode|TestEncrypt' + - name: Run benchmark tests + run: go test -v -run=^$ -benchtime 0.5s -bench=. ./... + - name: Upload coverage to Codecov + if: ${{ !cancelled() }} + uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0 + env: + OS: ${{ matrix.os }} + GO: ${{ matrix.go }} + with: + files: coverage.txt,coverage-root.txt + env_vars: OS,GO + token: ${{ secrets.CODECOV_TOKEN }} + - name: Upload test report to Codecov + if: ${{ !cancelled() }} + uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0 + env: + OS: ${{ matrix.os }} + GO: ${{ matrix.go }} + with: + report_type: test_results + name: Unit tests + files: report.xml,report_root.xml + env_vars: OS,GO + token: ${{ secrets.CODECOV_TOKEN }} diff --git a/third_party/quic-go/.gitignore b/third_party/quic-go/.gitignore new file mode 100644 index 0000000..60571ed --- /dev/null +++ b/third_party/quic-go/.gitignore @@ -0,0 +1,20 @@ +debug +debug.test +main +mockgen_tmp.go +*.qtr +*.qlog +*.sqlog +*.txt +race.[0-9]* + +fuzzing/*/*.zip +fuzzing/*/coverprofile +fuzzing/*/crashers +fuzzing/*/sonarprofile +fuzzing/*/suppressions +fuzzing/*/corpus/ + +**/testdata/fuzz/ + +gomock_reflect_*/ diff --git a/third_party/quic-go/.golangci.yml b/third_party/quic-go/.golangci.yml new file mode 100644 index 0000000..82bf4fc --- /dev/null +++ b/third_party/quic-go/.golangci.yml @@ -0,0 +1,99 @@ +version: "2" +linters: + default: none + enable: + - asciicheck + - copyloopvar + - depguard + - exhaustive + - govet + - ineffassign + - misspell + - nolintlint + - prealloc + - staticcheck + - unconvert + - unparam + - unused + - usetesting + settings: + depguard: + rules: + random: + deny: + - pkg: "math/rand$" + desc: use math/rand/v2 + - pkg: "golang.org/x/exp/rand" + desc: use math/rand/v2 + quicvarint: + list-mode: strict + files: + - '**/github.com/quic-go/quic-go/quicvarint/*' + - '!$test' + allow: + - $gostd + rsa: + list-mode: original + deny: + - pkg: crypto/rsa + desc: "use crypto/ed25519 instead" + ginkgo: + list-mode: original + deny: + - pkg: github.com/onsi/ginkgo + desc: "use standard Go tests" + - pkg: github.com/onsi/ginkgo/v2 + desc: "use standard Go tests" + - pkg: github.com/onsi/gomega + desc: "use standard Go tests" + http3-internal: + list-mode: lax + files: + - '**/http3/**' + deny: + - pkg: 'github.com/quic-go/quic-go/internal' + desc: 'no dependency on quic-go/internal' + misspell: + ignore-rules: + - ect + # see https://github.com/ldez/usetesting/issues/10 + usetesting: + context-background: false + context-todo: false + exclusions: + generated: lax + presets: + - comments + - common-false-positives + - legacy + - std-error-handling + rules: + - linters: + - depguard + path: internal/qtls + - linters: + - exhaustive + - prealloc + - unparam + path: _test\.go + - linters: + - staticcheck + path: _test\.go + text: 'SA1029:' # inappropriate key in call to context.WithValue + paths: + - internal/handshake/cipher_suite.go + - third_party$ + - builtin$ + - examples$ +formatters: + enable: + - gofmt + - gofumpt + - goimports + exclusions: + generated: lax + paths: + - internal/handshake/cipher_suite.go + - third_party$ + - builtin$ + - examples$ diff --git a/third_party/quic-go/AUTOCAR_PATCHES.md b/third_party/quic-go/AUTOCAR_PATCHES.md new file mode 100644 index 0000000..0fe526a --- /dev/null +++ b/third_party/quic-go/AUTOCAR_PATCHES.md @@ -0,0 +1,28 @@ +# AutoCAR security hardening + +This directory is based on AutoCAR's pinned +`github.com/apernet/quic-go` pseudo-version +`v0.61.1-0.20260806010916-184d081eef3e` and remains licensed under the MIT +license in `LICENSE`. + +AutoCAR adds one narrow HTTP/3 server hook: `StreamAdmission` runs immediately +after a bidirectional stream is accepted, before a handler goroutine is started +or the first frame type is read. It can reject a stream or return a callback +that releases process-wide capacity when handling ends. + +The Hysteria server adapter uses this hook to apply its global handler budget +and first-byte deadline before `StreamDispatcher` peeks the frame type. Without +that ordering, a peer could open many streams and send an incomplete frame type +without entering Hysteria's normal TCP request handler. Its dispatcher handles +the TCP relay synchronously in quic-go's existing per-stream worker so the +release callback cannot run while the destination header or relay is still +active. + +The pre-handshake server path also cancels `ConnContext`, closes a just-created +qlog trace and releases the Initial packet if connection-ID generation fails. +This preserves admission-slot accounting even during a system randomness +failure before a connection object exists. + +The testdata helper generates an ephemeral ECDSA P-256 CA and leaf in a private +temporary directory at runtime. Fixed test private-key files are deliberately +excluded from the fork. diff --git a/third_party/quic-go/FIPS140.md b/third_party/quic-go/FIPS140.md new file mode 100644 index 0000000..7d7d92e --- /dev/null +++ b/third_party/quic-go/FIPS140.md @@ -0,0 +1,37 @@ +# FIPS 140-3 + +quic-go relies on the Go standard library for cryptography, including the Go Cryptographic Module described in [The FIPS 140-3 Go Cryptographic Module](https://go.dev/blog/fips140). quic-go does not seek separate FIPS 140-3 validation as a cryptographic module. This document explains how quic-go uses Go standard library cryptography for QUIC operations relevant to FIPS 140-3. + +Starting with quic-go v0.60, the behavior described here applies when built with Go 1.26 or newer. With older Go versions, quic-go still builds and runs as usual, without any attempt to meet FIPS 140 requirements. + +## QUIC operations relevant to FIPS 140-3 + +quic-go delegates the TLS 1.3 handshake, certificate handling, cipher suite selection, session tickets, and the TLS key schedule to `crypto/tls`. When Go's FIPS 140-3 mode is active, `crypto/tls` restricts the algorithms it negotiates. + +### Packet protection AEADs + +The main quic-go-specific FIPS-relevant operations are the AEADs protecting Handshake, 0-RTT, and 1-RTT packets. + +AES-GCM packet protection AEADs are constructed through the Go standard library's TLS 1.3 AES-GCM implementation. Today this uses `go:linkname` to call the unexported `crypto/tls.aeadAESGCMTLS13`, because the standard library does not yet expose a QUIC-specific constructor; see [golang/go#79219](https://github.com/golang/go/issues/79219). + +ChaCha20-Poly1305 is not used in Go's FIPS 140-3 mode. `crypto/tls` avoids that cipher suite during negotiation, and quic-go additionally guards its internal ChaCha20-Poly1305 path when FIPS 140-3 mode is enabled. + +### Header protection + +For Handshake, 0-RTT, and 1-RTT packets protected with AES cipher suites, header protection keys are derived with `crypto/hkdf` and the AES block operation uses `crypto/aes`. ChaCha20 header protection is tied to the ChaCha20-Poly1305 cipher suite and is not reachable in FIPS 140-3 mode. + +### Address validation tokens + +quic-go encrypts the address validation tokens it sends in Retry packets and NEW_TOKEN frames. These are not TLS session tickets (those are handled by `crypto/tls`); they carry server-defined state such as the client address, timestamp, RTT information, and Retry connection IDs. + +Token-protection keys are derived with `crypto/hkdf`, AES is used via `crypto/aes`, and the token AEAD is constructed with `cipher.NewGCMWithRandomNonce`, keeping token encryption on standard library primitives. + +## QUIC operations not relevant to FIPS 140-3 + +### Initial packet protection + +Initial packet protection (including Initial header protection) is not treated as FIPS 140-relevant confidentiality protection: the Initial secrets are derived from constants in RFC 9001 and the packet's destination connection ID, so any observer can derive the same keys. quic-go therefore disables strict FIPS 140 enforcement around Initial packet construction in Go 1.26 FIPS 140-3 mode. See the IETF QUIC mailing list discussion at . + +### Retry packet integrity tag + +RFC 9001 defines the Retry packet integrity tag using fixed keys and nonces. It guards against accidental corruption and casual injection but does not encrypt packet contents. quic-go treats it as outside the FIPS 140 scope and disables strict FIPS 140 enforcement for that AEAD construction in Go 1.26 FIPS 140-3 mode. diff --git a/third_party/quic-go/FUZZING.md b/third_party/quic-go/FUZZING.md new file mode 100644 index 0000000..e7500aa --- /dev/null +++ b/third_party/quic-go/FUZZING.md @@ -0,0 +1,59 @@ +# Fuzzing + +[![Documentation](https://img.shields.io/badge/OSS--Fuzz-Introspector-red?style=flat)](https://introspector.oss-fuzz.com/project-profile?project=quic-go) +[![ClusterFuzz coverage](https://img.shields.io/codecov/c/github/quic-go/quic-go/master.svg?flag=clusterfuzz&label=ClusterFuzz%20coverage&logo=codecov&logoColor=white&style=flat)](https://app.codecov.io/gh/quic-go/quic-go?flags%5B0%5D=clusterfuzz) +[![ClusterFuzz Lite Batch coverage](https://img.shields.io/codecov/c/github/quic-go/quic-go/master.svg?flag=clusterfuzz-lite-batch&label=ClusterFuzz%20Lite%20Batch%20coverage&logo=codecov&logoColor=white&style=flat)](https://app.codecov.io/gh/quic-go/quic-go?flags%5B0%5D=clusterfuzz-lite-batch) + +Run the commands below from a local [`google/oss-fuzz`](https://github.com/google/oss-fuzz) checkout. +Fuzz target names match the binary names listed in `oss-fuzz.sh` (for example, `frame_fuzzer_v2`). + +Update the base images: +```sh +python3 infra/helper.py pull_images +``` + +## Running fuzzers locally + +The following steps run a single fuzz target and then open its line-by-line coverage in `go tool cover`. + +```sh +export DOCKER_DEFAULT_PLATFORM=linux/amd64 +export FUZZ_TARGET= +export CORPUS_DIR=corpus/$FUZZ_TARGET + +mkdir -p "$CORPUS_DIR" + +python3 infra/helper.py build_image --no-pull quic-go +python3 infra/helper.py build_fuzzers --sanitizer address quic-go +python3 infra/helper.py run_fuzzer --corpus-dir="$CORPUS_DIR" quic-go "$FUZZ_TARGET" +``` + +Leave `run_fuzzer` running for a while to build up a corpus. It unpacks the seed corpus zip into the corpus directory and appends new entries as it discovers them. + +```sh +python3 infra/helper.py build_fuzzers --sanitizer coverage quic-go +python3 infra/helper.py coverage --no-serve --fuzz-target "$FUZZ_TARGET" --corpus-dir="$CORPUS_DIR" quic-go +sed "s#^/out/#$(pwd)/build/out/quic-go/#" build/out/quic-go/fuzz.cov > "/tmp/quic-go-$FUZZ_TARGET.coverprofile" +go tool cover -html="/tmp/quic-go-$FUZZ_TARGET.coverprofile" +``` + +The `sed` command rewrites the container paths in `fuzz.cov` so that `go tool cover` can locate the source files in the local checkout. + +To produce a coverage report against a modified local source tree, mount the local checkout when building the coverage fuzzers, the same way you would for reproducers: + +```sh +python3 infra/helper.py build_fuzzers --sanitizer coverage --mount_path /root/go/src/github.com/apernet/quic-go quic-go +``` + +## Reproducing an OSS-Fuzz testcase + +Download the reproducer file from the OSS-Fuzz report. To test a local fix, rebuild the fuzzers with the modified quic-go checkout mounted at the path expected by `oss-fuzz.sh`: + +```sh +export DOCKER_DEFAULT_PLATFORM=linux/amd64 +export FUZZ_TARGET= + +python3 infra/helper.py build_image --no-pull quic-go +python3 infra/helper.py build_fuzzers --sanitizer address --mount_path /root/go/src/github.com/apernet/quic-go quic-go +python3 infra/helper.py reproduce quic-go "$FUZZ_TARGET" +``` diff --git a/third_party/quic-go/LICENSE b/third_party/quic-go/LICENSE new file mode 100644 index 0000000..51378be --- /dev/null +++ b/third_party/quic-go/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2016 the quic-go authors & Google, Inc. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/third_party/quic-go/README.md b/third_party/quic-go/README.md new file mode 100644 index 0000000..22056d9 --- /dev/null +++ b/third_party/quic-go/README.md @@ -0,0 +1,65 @@ +
+ +
+ +# A QUIC implementation in pure Go + + +[![Documentation](https://img.shields.io/badge/docs-quic--go.net-red?style=flat)](https://quic-go.net/docs/) +[![PkgGoDev](https://pkg.go.dev/badge/github.com/apernet/quic-go)](https://pkg.go.dev/github.com/apernet/quic-go) +[![Code Coverage](https://img.shields.io/codecov/c/github/quic-go/quic-go/master.svg?style=flat-square)](https://codecov.io/gh/quic-go/quic-go/) +[![Fuzzing Status](https://oss-fuzz-build-logs.storage.googleapis.com/badges/quic-go.svg)](https://issues.oss-fuzz.com/issues?q=quic-go) + +quic-go is an implementation of the QUIC protocol ([RFC 9000](https://datatracker.ietf.org/doc/html/rfc9000), [RFC 9001](https://datatracker.ietf.org/doc/html/rfc9001), [RFC 9002](https://datatracker.ietf.org/doc/html/rfc9002)) in Go. It has support for HTTP/3 ([RFC 9114](https://datatracker.ietf.org/doc/html/rfc9114)), including QPACK ([RFC 9204](https://datatracker.ietf.org/doc/html/rfc9204)) and HTTP Datagrams ([RFC 9297](https://datatracker.ietf.org/doc/html/rfc9297)). + +In addition to these base RFCs, it also implements the following RFCs: + +* Unreliable Datagram Extension ([RFC 9221](https://datatracker.ietf.org/doc/html/rfc9221)) +* Datagram Packetization Layer Path MTU Discovery (DPLPMTUD, [RFC 8899](https://datatracker.ietf.org/doc/html/rfc8899)) +* QUIC Version 2 ([RFC 9369](https://datatracker.ietf.org/doc/html/rfc9369)) +* QUIC Event Logging using qlog ([draft-ietf-quic-qlog-main-schema](https://datatracker.ietf.org/doc/draft-ietf-quic-qlog-main-schema/) and [draft-ietf-quic-qlog-quic-events](https://datatracker.ietf.org/doc/draft-ietf-quic-qlog-quic-events/)) +* QUIC Stream Resets with Partial Delivery ([draft-ietf-quic-reliable-stream-reset-07](https://datatracker.ietf.org/doc/html/draft-ietf-quic-reliable-stream-reset-07) and [draft-ietf-quic-reliable-stream-reset-09](https://datatracker.ietf.org/doc/html/draft-ietf-quic-reliable-stream-reset-09)) + +Support for WebTransport over HTTP/3 ([draft-ietf-webtrans-http3](https://datatracker.ietf.org/doc/draft-ietf-webtrans-http3/)) is implemented in [webtransport-go](https://github.com/quic-go/webtransport-go). + +Detailed documentation can be found on [quic-go.net](https://quic-go.net/docs/). + +## FIPS 140-3 + +Starting with v0.60, quic-go supports use in FIPS 140-3 environments when built with Go 1.26 or newer, using Go standard library cryptography for the QUIC code paths relevant in FIPS mode; see [FIPS140.md](FIPS140.md) for details. + +## Projects using quic-go + +| Project | Description | Stars | +| ---------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------- | +| [AdGuardHome](https://github.com/AdguardTeam/AdGuardHome) | Free and open source, powerful network-wide ads & trackers blocking DNS server. | ![GitHub Repo stars](https://img.shields.io/github/stars/AdguardTeam/AdGuardHome?style=flat-square) | +| [algernon](https://github.com/xyproto/algernon) | Small self-contained pure-Go web server with Lua, Markdown, HTTP/2, QUIC, Redis and PostgreSQL support | ![GitHub Repo stars](https://img.shields.io/github/stars/xyproto/algernon?style=flat-square) | +| [caddy](https://github.com/caddyserver/caddy/) | Fast, multi-platform web server with automatic HTTPS | ![GitHub Repo stars](https://img.shields.io/github/stars/caddyserver/caddy?style=flat-square) | +| [cloudflared](https://github.com/cloudflare/cloudflared) | A tunneling daemon that proxies traffic from the Cloudflare network to your origins | ![GitHub Repo stars](https://img.shields.io/github/stars/cloudflare/cloudflared?style=flat-square) | +| [frp](https://github.com/fatedier/frp) | A fast reverse proxy to help you expose a local server behind a NAT or firewall to the internet | ![GitHub Repo stars](https://img.shields.io/github/stars/fatedier/frp?style=flat-square) | +| [go-libp2p](https://github.com/libp2p/go-libp2p) | libp2p implementation in Go, powering [Kubo](https://github.com/ipfs/kubo) (IPFS) and [Lotus](https://github.com/filecoin-project/lotus) (Filecoin), among others | ![GitHub Repo stars](https://img.shields.io/github/stars/libp2p/go-libp2p?style=flat-square) | +| [gost](https://github.com/go-gost/gost) | A simple security tunnel written in Go | ![GitHub Repo stars](https://img.shields.io/github/stars/go-gost/gost?style=flat-square) | +| [Hysteria](https://github.com/apernet/hysteria) | A powerful, lightning fast and censorship resistant proxy | ![GitHub Repo stars](https://img.shields.io/github/stars/apernet/hysteria?style=flat-square) | +| [Mercure](https://github.com/dunglas/mercure) | An open, easy, fast, reliable and battery-efficient solution for real-time communications | ![GitHub Repo stars](https://img.shields.io/github/stars/dunglas/mercure?style=flat-square) | +| [nodepass](https://github.com/NodePassProject/nodepass) | A secure, efficient TCP/UDP tunneling solution that delivers fast, reliable access across network restrictions using pre-established TCP/QUIC/WebSocket or HTTP/2 connections. | ![GitHub Repo stars](https://img.shields.io/github/stars/NodePassProject/nodepass?style=flat-square) | +| [OONI Probe](https://github.com/ooni/probe-cli) | Next generation OONI Probe. Library and CLI tool. | ![GitHub Repo stars](https://img.shields.io/github/stars/ooni/probe-cli?style=flat-square) | +| [reverst](https://github.com/flipt-io/reverst) | Reverse Tunnels in Go over HTTP/3 and QUIC | ![GitHub Repo stars](https://img.shields.io/github/stars/flipt-io/reverst?style=flat-square) | +| [RoadRunner](https://github.com/roadrunner-server/roadrunner) | High-performance PHP application server, process manager written in Go and powered with plugins | ![GitHub Repo stars](https://img.shields.io/github/stars/roadrunner-server/roadrunner?style=flat-square) | +| [syncthing](https://github.com/syncthing/syncthing/) | Open Source Continuous File Synchronization | ![GitHub Repo stars](https://img.shields.io/github/stars/syncthing/syncthing?style=flat-square) | +| [traefik](https://github.com/traefik/traefik) | The Cloud Native Application Proxy | ![GitHub Repo stars](https://img.shields.io/github/stars/traefik/traefik?style=flat-square) | +| [v2ray-core](https://github.com/v2fly/v2ray-core) | A platform for building proxies to bypass network restrictions | ![GitHub Repo stars](https://img.shields.io/github/stars/v2fly/v2ray-core?style=flat-square) | +| [YoMo](https://github.com/yomorun/yomo) | Streaming Serverless Framework for Geo-distributed System | ![GitHub Repo stars](https://img.shields.io/github/stars/yomorun/yomo?style=flat-square) | + +If you'd like to see your project added to this list, please send us a PR. + +## Release Policy + +quic-go always aims to support the latest two Go releases. + +## Contributing + +We are always happy to welcome new contributors! We have a number of self-contained issues that are suitable for first-time contributors, they are tagged with [help wanted](https://github.com/apernet/quic-go/issues?q=is%3Aissue+is%3Aopen+label%3A%22help+wanted%22). If you have any questions, please feel free to reach out by opening an issue or leaving a comment. + +## License + +The code is licensed under the MIT license. The logo and brand assets are excluded from the MIT license. See [assets/LICENSE.md](https://github.com/apernet/quic-go/tree/master/assets/LICENSE.md) for the full usage policy and details. diff --git a/third_party/quic-go/SECURITY.md b/third_party/quic-go/SECURITY.md new file mode 100644 index 0000000..d8f0570 --- /dev/null +++ b/third_party/quic-go/SECURITY.md @@ -0,0 +1,14 @@ +# Security Policy + +quic-go is an implementation of the QUIC protocol and related standards. No software is perfect, and we take reports of potential security issues very seriously. + +## Reporting a Vulnerability + +If you discover a vulnerability that could affect production deployments (e.g., a remotely exploitable issue), please report it [**privately**](https://github.com/apernet/quic-go/security/advisories/new). +Please **DO NOT file a public issue** for exploitable vulnerabilities. + +If the issue is theoretical, non-exploitable, or related to an experimental feature, you may discuss it openly by filing a regular issue. + +## Reporting a non-security bug + +For bugs, feature requests, or other non-security concerns, please open a GitHub [issue](https://github.com/apernet/quic-go/issues/new). diff --git a/third_party/quic-go/assets/LICENSE.md b/third_party/quic-go/assets/LICENSE.md new file mode 100644 index 0000000..672a20f --- /dev/null +++ b/third_party/quic-go/assets/LICENSE.md @@ -0,0 +1,27 @@ +# quic-go Logo and Trademark Usage Policy + +## Exception to Main License + +The files in this directory (collectively, "Brand Assets") are **excluded** from the quic-go project's main MIT License. These assets are protected by copyright and trademark laws. + +## Permitted Use + +You are granted a limited, non-exclusive license to use these Brand Assets solely for the following purposes: + +- **Editorial and Press:** You may use the Brand Assets in blog posts, news articles, video reviews, and public presentations that discuss, review, or reference the quic-go project. + +- **Reference:** You may use the Brand Assets to indicate your project's compatibility with or dependence on quic-go (e.g., "Powered by quic-go"). + +## Restricted Use + +You may NOT: + +- **Modification:** Modify the Brand Assets in any way (including changing colors, aspect ratio, or obscuring the image). Resizing the image while maintaining the original aspect ratio is permitted. + +- **No Branding:** Use the Brand Assets as the logo, icon, or mascot for your own project, product, or service. + +- **No Endorsement:** Use the Brand Assets in a way that suggests your project is officially sponsored by, endorsed by, or affiliated with the quic-go maintainers. + +## Termination + +The quic-go project reserves the right to revoke this authorization at any time if usage is found to be confusing, misleading, or detrimental to the project's reputation. diff --git a/third_party/quic-go/assets/logo.svg b/third_party/quic-go/assets/logo.svg new file mode 100644 index 0000000..ecf2bc0 --- /dev/null +++ b/third_party/quic-go/assets/logo.svg @@ -0,0 +1,330 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/third_party/quic-go/assets/quic-go-logo.png b/third_party/quic-go/assets/quic-go-logo.png new file mode 100644 index 0000000000000000000000000000000000000000..6be4ed43184346bddc76d6ff258eddb05488bb50 GIT binary patch literal 105123 zcmeFY_g7Qf^FEH{%2iYZ1f*-EBOtwlC4ka95}Nc9nzTsAMlVuBN9irpP(mmwz4sPM zqy`89siE@;-q)4;{sX?hefL`7pMNf)7Pf$dhcvFTuUs`9EnOid&K5)v2!zMR z0qkOC;$*?&=xmj|E=fy7bf4%Y@Tu0@l(lI{LhQ)F+2)a?fP{XoeeR7H8}%aL>SZRE z=rUTz!^>FWn~lgbA6P#sdrh(zQFa+nkFWCIF?M@@tDJ=b{?O1I}N_jbg_K-<`B;SbP7zzyMZBv(i1F zz>%$?F*j-Js6;%}Pt!S!gPWKKjpBeWbF8OzVmt40(~vW0#)4(Do&gnrSC-CFUmMHm@`6wSuZ#YnaMTnDnX$Uj8Y)Rft^2M5%nwdcM}tRyUxnO(qfWS znacr_vyCzV@T;9ljljNo1Du}#g@ZNpfjgdLU2}vHO?90ZWA z54zgpO+=JJ6sz%u2YR9bTQp^HE(mH4*#EmR$W9BeeFE|U>F89+rQb)Jvu77MKumdI z;okwCr3tpB>?>U1;LpIOk)ZZV1Le;-CTv311eGRcDJJxEbad{lI9~P%1Ix?H4Wn*f{qvJV z?|v`Ko7Pz;*Xnz#SIDFu59;|^rt&Y&W|WunK4i$^;ei%bJUA5qh6RImfMgv(PFZKq zL$&Av)E&w|oUtr9AYgMT%Jx=h3opP_@kId{8J`_r&Ipxb2#1Egw&ks=^nB&TZLJ9) z@L1Y0Z11QD0Prmld73OP2bXyY_Bl47DOxPOSN;U(tJW#rn3uQwI_u$uIs5;8zHIcm z20TYvEURv0FKw4}jIy<|snjzn1@2Cq%6E}l@HujD7L1b*`Xql}biSpdph^W}3vKsz zN`u963Y*$}s`{8!E()EDNk~~yv>kd%Ps3nJKEyviXTY)fMg>yG-4S&+c>%6(YF9Sf zjW=$fRn4x>$SBAu_@Qa2W1y$62h!J9HzIsiw`1=f7H}yjIeH}S>1Y=ug=l2v6b#Po z#!enHs@9cMv8o^W+=f-^1GSj9#!O3eiu|S$BzI zGVYR~pzNo-?Ana!DYXI>V?RT>FE!0;h~pJhI@dniD!x!S_Gn5#>+5Sq`H40@!asSl zIm2Lq&>uYrC0$e6x=Wol)tqlHJ+Ni)Ec)d1dEjkrq1h8lb%hqu)L@EOjTqT34?uw# z2_vTRmD%kU#FRXPy@YOd+LV5_`v!Djed-}MA51xEgofn!% zg1O6)xsOX)H4`v}oScsf6-nZ9;o}@swX|Ntxs}>1$LD}wQ^n8E3445jUf^?J1ZyLe0)_)DTT^I z-RLr#vGq?bTYJ0*y@vYgqq@MbYmQ&cks3?cvq^&(SMfla00S0~F)N_s>kPUSM2Z@k zWKh_U%F9tciwfp2Jj#+lx!YO2DtcOikUg>OR8-`QVB}<%Nioz(sn0EN!dUvf*_#nn zRrd`*W9-=geB4fUQ9<4XE#8fAxbxRh6#BB&h2-DP&bP*sSj9J@a)|-Q_W@5!j7IaxYWDqO0I zd;0W{HA~CI_Pt3Vb1|}Due{!V0779UQfcTd5)OBD%iH_wkFfsf8RwpXeaRhZXnI&3 z99{K&bUZpGneF62*~>Y1J++BS30&_tF?FlU*t%71&*2 ztl>yJQfeCIF{Ninp)oj^T)_k=+**#BnkH$Ss_@$ssb%(pP!Da8qKGD+2q_HHH^z^3 zb#^X&%D7JXN2+hZ;g-@wL>a?=#bbu708X{~lb>awz6ZY)=VkN#cbx`vY>NFb8(mf2 zf`yTz;)t}7Z3P9c-VxB4psOFWYediH%syjS2rxK|)pk?z{t*$;q|$+v^6;D-UO%Z7 zdBX&Jmm#I+Rpoc!zV|MUM8lS=^DroAAPnaos=>={;gOk_mnW!2bK&JuUj5dS0>W?) z0a%5;v)v2Z$|}ubuUH+7(b<(e+FB2`vFA@NCC)I=-W+(JYV22IhJP={w|EeEjXkP{ z)j|8KY){pG4~C?xt2$7S`Hy=1=NZqPF7WH!K(igcoGpef(*Is=l6;x|}YXDD~j=fp+vDscqC0%3Oadr2^ zlz|7Lx%X`|zp6aLyKQX4uJdrKe!n0i^Gddjas~+KQz9#6sG-3xGW;LC(5L^_%MFBK zL?qB$;S1`QF}l94??5;Fwo8`Wq!2zvyDQDBv@xieSU(epHrCQOZT35xTyb^RRAAK3 z;EiuJ3M~k$+xyVAcF%JW$1r8+>2Blfl=8SBB$CSRU||=BY&;W0RApp14!Nt-#l`1_ zc6CY*gN~jh-}$4CZ<`-~TT9rqgK7mxDLCytv7RZXCTrn6v>dKlI#1S@1u!Z{Z^D{gtuL_9Ir%8HEmiCVO?-}$-jytC)zP$( z5}DK(q=j3fxU>Lnf@tQG@9^Y!8W3aF5dBE{W};Ll=son-kfhsJp6buVn9S@TxN9EIgp+4j@ruC zlL@eqkF#U%VfGdHXt7*_(S3#bFR>Pu>gwt?p;s@U$Q|w9iyb1G+H3xN#hXlQsdn7` ziX6B2$)+9!;0_mad4o4(u@jnp?*SmZ0WeJU_RWj7Z&t{OYS@hsyE|Y>sSOW@u^A>-dtnZr7h=`t;TUc7Maz?Gcvt^J^w+xU_Rea;N>^#^((WggQWTfXR7%wLc z3|&x}mPi-b{_sZA%11Y0)a9osS-wp>VG3E+mF+n9V=l1?#*s0ri*>QS$9g^s7LAnN zsGLEr%Ok8(ViH^xeLu}u_dZ28R-t0*J=Jx311uJ^1NS&^0xc5_6}@W$f?ucm`BEvS zZ1NvSbYWpWM}sC}mau9?|8i+eKEyqHmv;Q27Ut5dMikWL?ErdxI2I6;Ea}Fe zUn^+) zp6OcmQfqtC@@{CtsO3pNmxDlUisyR2S_;y>SD#C#&(l`x3y2Ve71?C)6XR+rSAfrC zje2K4siG^6mV3i9fy_DYN<3oQ(XgtTw(ae0E)YM#TK+fuf1Ri5{_F&-$V3KXBG^tS zw|ie)dXWr*rT6&F3JOs7tz2>8F?nmL<4-8gdvI~+uZ8?9kGP!$Q$r z^jdO2SU-yXE6_899U63u3ET_HGbNjV$y#u5cik$)2mbg`kn&>T0*7S0BU>Q|{>3DA z)c&$mH4AfGRkO!rV?w>@3i?bmi!a<80Em;1ym__|@bV2pf6T$L_n`A}*mvj|tH=v6 z-+3=ydz;?VC*Ys!!?2-P-3A7aC9{gE2?CO#+lh{@0z~1W<$`OKfVaq_dI57(WGorvz^v3DCAAF4yV++phcMz8!-%*Jo+IUz zm6eyX-M9W#F%j^}Z#PE{dil+=fsiRstX72VZJ?^RM%aj7QZ(xJqm&Mo8VXg#>r4(t zmll@YW+

X~@2__43`9B*3*Ru$Rg2G7nX}96iqXs-@Ta)_I|^Ti zvAo^}nL6k24qNs3?FU?DAs?$Vbe*e{cTg(k;Noyty&Zo~-%bww;1b)f5go0a{Pr%f zsU7bp*rAi1P&uVDC%3j!YU65PY}A;)?xTzm$Jhi3L!P!uh5XOHejCn+Ye>TgCSVr^ zFv$B+UqUut-g_nqr;@uJc4wXBLPgyFs7Rrid>O$atP{AE620a_%tf*tI9is~8vhy2 zJ-Uv{id9z|Gs;1u_mX9dMvAcR#Kv#D@zjS&$xk`5r&NxvD0B+}_{$JSBVX&k?uc~ov?&_Y{}ITa zg5kGzBK?i4<8F-4kD1Hx?63I;DB{9HgVXi$OS(P{93DTW1xEGE z<{^Jgmb^V%k7Akd+7l5{l`DBC^qSQ65GXLl@ELBKY`pxl1deYk!eG>3cEnXi0 zu)JbKXUkqcg>$Z3auU@@b}sH(bVeHN@dESZStR|8)IbZ@9F=!F;qc}v6VVH%%{p$M zgg^~s$bY*gVR6-a%DlzP{i@wP{zvgmG3h6+S^SUAu#8#(CayG)rOy!$Rq{`FS}SyQ z8Kw9ozZJ$?^vQVtT5IfHAI^HBkiVcJ<@*Q)b?g4vuyvf`u8g7x0&K1i$GTY7$?!lX z4~q^WXBO}}+B+}|<#S+h**)hp0#+!@%Cf5(t$j)%IW$N<$yxU3F84I;y*{gO#qm^6 z=RR-wHV5(VskXt&&LGK8h?(+-VwRhysbv=OxA%rYNbK~>mrUtn<;{!EsYV(aKZce6 zm_v%wKY-Pm%<)T=R7oyp*0lGcd}Q-g0u?c{_l5t_rCdur`?5FGb-z3|AiQ)jW`{mF zKO^381j5CI-|cAeuXybebB67y$FYrKN0$ef3EtgEx2{pc@h!e3J6ZZ|e(tEHPYhl| z7F%tg(iTTB$JW8Y!D%;ccLK3OWs%5La4Of3@G%SK)7w~(wQ-<(J1kgrl>;F40XC9@ zPDgs|j%t@kX@lzg#+=Bm?~5(T)UsRfyb-0RAY=pYRgN!UM})jV1qmU|Q&YHcS*nDS z^Y0ZYBvql%&#SyR5v(08;b#?AU1oOb-`VD%$G-0T`o5N9x%y_Q<$0`$)bTckEnIs3 zi6;JPMzn@)KU}}nXm4gKpzkp?j9rb%%5fW7n**esYxu!$!MnPG&3cSzN~^QDUW<3y(w=R@PS2(B%G<6aV?Aws_xZ zUO!0U%;5@M*JzzkMWMAoW}5C{QotiOL<%oWz#^%vtZX&`$dT0DOh&$x0U z)y0=6!5%e+mC(66^X zh7WJ8Wj7^L(F?mFzh(g0by34@wwo-6mN>jU1j*1y&dp0~XdOD?wMmC&%W>O4OTfR% zt|j6oRXlu*q6yw2nR^8;;x}o=o&&w# z+3FeBCmW-W*uwW)`1Edu3L#HEmq|UpL;BhvwAGBja)71ub z;B`EJOSuOeO-Wv}+s3V}NKMUWsk@Ej1aw7Re;lqzGdgu&F@-90(@8P)!sEeNf8$S< zW8X@aSGe{Z?eDZEv^})x3D$})P&9c}el2xIz9rRF#rU@A&4p#*lK^gVvdEwyAs&-9 zF+U$~B?ny53=fr9JXR zUdoKO++(1vx>mdQymfcuvAIEMAGgj>#`@p77p0i{8vmpg>`|2O(x`DWgAZacH*4oc zdjT0&M)tT5+>hVOpsV_I=ktreq>4&mgy8A=iJ#iP^-UFSd^Wn8>Nbu0u3=X!{|(wh zBdBJdqIHMCyRfqpV`8qRzo%Y2a@GTZlB9|*I7R=*|G>lyU?0iOlr9!$ke2Kss)eZ$ zcNObM8}QZsT>j?u!EvyoD#QB3X$j=AF!g0?!$Tw2o&-g-o*gq}wM^^1;OgmcbRzzi z5&t$Ri7eK0EFgrNuVZF=t*e@tgY6Jp*mwKCb0q&~j{B8-GOz51GKPSFSH?BZeyzzhnf4T0AuI@l1cl7;t_|3> zQz4{@Tn%{5&{NZoEdRRJP(TVrG5sG4f5M#i`1T~>_Q6kgvAy6Ix#=?1<8$uD2OYJ= zUnRjKMU9VUq%rC9SkMp-zD*rbj|&baycU*DU``Dap)uGF83T>DbnIp<)myMaXg8{y zoCPqqJsQ;5YId(P!;BmU0y48ozig|G2^ZB(65?Uu7UsxI72vDz$r0k=<0zL&ridhP zp<)|fv15x4`Z7AJ2gK3=ElDn|^CfXsdkiQp`U&a!*&@p9K+-Vrsti&pXl(l3h@1bJ z!K;T1gge0qDi77*Df~9(6t!CHF)6avzATkzmDfxA1}8P&RgUNB)N6EoRw~qKs{K`* z6c$-hTb^+rJfD`SKN@>?W7np z1iFdmaPc88aq8%Vfv=-|j|?xpdCU{0Km}th1A|i=#nk_#9^}7$)dOhPn^bZS7By~| zS0f9i`>mgWVygY2PXnr8t6UT*fynNC(pA_k6C1NTP&ADIQ$hGoW>lhjhpYowL^9E( zsr5MqV>1K!lO0hJE_Li3GhRPKCv)GPR0r$nX}(e46k++L)ReT0cW>b>Y0D{-H`b7oCq)f+ZR9sK5~W+T+yh6WM$n z4@1r$UAl&PWo~;e*m9`%m@?S!%eDTceyS#C<_;4XT539McxdD{3g>QYsu7)-kdQU) zP|eBt=yQ4b4Q)Mj2Q3{PJ*hMWeQgVEjqJ7|m0;!o&LHy4FFmDi^ob2zbk7o)sbd}{ zKTXaNI~z!bE)?l&VBBi2@HSf5G?MrUSx@Wt6?aWGs31`8TcpVKG~aUwo|BHlH;N$| zIRJcZ!$4yV4UJP1P5gNW!TBIjHc0%p7l4QeB=9?p(m2mRjU0D*1_oKxe3o}6p}WZa zTE5xdPoqux+Uiv+;WXIDBGBU}V4TCcuTG&*t_!r6_;6=s9*6dOZFlzciiYRT@_;m9 ze9iY-T3Dc;skKm?X}v|)DdLFAswj8QiiO7Q-BlmWVbi8uxgu9Y=Q{1eucGm(%bk$SstBiR0u z?h~!g$}%_7{7=;ZB=LJTrViEfx*06L;$jryjCQuhEFy#%BtD8!0hEqTW*E}z8l3Ts zGb`CGO7^SLRG}=5K7o5)tf_8Qj-A6Csp4+y^zo4_s@mF%KbWX+$4gk=F5ME41ur}M z8r!8D_=tR{mgVhx2=`@fQqsz1wvTp~e+FCb7Ujm?>Ee2|vg$*LCC=_O{GGm&MjYcjZip-<^ugQHF@K^$5ekvc_HicJ`0}OE{n}1LDMgIC0vcYzm9tesLj`kn($W58@ z>hS3E7A!8q)VZDd+YN1y?0CUTYqHj2Th}j0Mx*6A4O=;+2Y4nXg1l9sHP{L0q{ zdr1nEkpsvxbE&$#>~RAR-KP_M4$Dr@>VpCp32wr`*8+%PB!>ZMd~6GU?~ ziF)p0k~QhJu_;J5+@9k)vHo`zg7T2JzI09{dzOa4R2bAgRD%0oQ79>IZdU z(tWE|2ILvb;-PUuZwp?hNZ8{(${UMujK)QeMju3ooU_ng&`9oIM|G*Ec(pEab#kZ6 zXab)lW$rli^*hB~nTEZvI0Bn{DOZuRiC#@hs|-<=c;Sex$wvG<2wcahG8m^Fu=|XR zd%V;`GaX~}_KlB$_tbhoovxv`>oVbA`L~6OU!ezSBH7NffZWsBL=Cco zvok>Z=a>&L-tpu(bbmdGyW_N88dN~iNxLLB>H_JvA@j~i42Xz+ITF$;;>vqNz-A`% zVEaYINYd%U4s6$tM3S%7yQvszUcaqg9+ab{uz9^n*~0W2OsY%~8eWFqzfN&j>Durj z(RDBDJ6*NiiXbk;MlmkB^KcYA)cs0mKFv4zRLn(Pbc=x(MVc&nzh3?<`^@VLi_|zM zJ$G>OvsC2@-VguHO^AGg^$ zS!~7tdSHL(t;mG5n4NKyq2~ouU8|d8;B#)CbE6Ysp4yt!R~2Jq-5hDNtpqW79B-B> z#`Bge&t!T%r`0vP-F;;DV=mC)RnCdDaVW8*%JEi|?zNzn_M<+(YbPlbG1iXlZTXK9 z-)uyqLds7sg@SHVLv`vVmOuY%hXk@gZ2jyRHIYAsY~XFqnH_x@il|u6cv>WNcjEr$ ztr(pH$+%q=-sjk#crESTBiyjEEV4?;uKln%tGTG8;Yb*hibpd!R!`!%=}SgB?1ot9 zu^SC^+rwdv#e6rlwa|yr^?GJY({DQD8goO=*l54STSIolv0DrPK!$}?savlw8)9`5 zypd5dx~_M;JN}bvI0l4WWo`bC3PViU3Clg8%8{L-(uP#n$qx4*RE5{1PTu)xRTA=k zM@5V0+ao%0P;b9BmWi2ykN|_ILisTIvW1?B3i;ONW()wDl)1r#@pncIOhiq zAMomVYsb|uq4{<{S7&6+*!Z%EhjQ~Y;J&F(!lTEZxZz(;SB;FcSD(}S4SsV~4&?AC zRqvd5I_rz-(zoFEzl#ScllJL(vW%oUx&<>_a@>K!u#{hWEKQ2Sb@(|A`@@q55kag4 zFm|p)zcH4-2=X%S!g|`mErKjL^nw-ASW?h4Fm4@=ht*LP4)Huk}VQpTsgsJW4fG+Guw@b|q@fFCBUw!Aa--XC(Rd8UMLM2VM`VRkINn$v9J<6DaJH7(O+_3s98J@2fXD^P;4r^x*!b z`A79cpuhfsf`o1~Ef(7bf~xsfJw~%b<)(MOeN30$j-{gX@xWxu=qxbNs-sU8@-qrC zc5jB7@A@h-)|!}X&J+djJ}NvosmB-9&$iOFeSA2y$u1*vw#>_#qOrWGwOa%>O#VJk zbsIS!9BdLo>84Y^wJ?V}OJcAdOw;Ai&aZU;w9H--F*SjP&t$afMo9F<-~F( z9-``2o8iK4Mpi2vJnSh)Ki%|Xk09YuM9S|vOBowHdvmHDgoIsd*Kc&4FG$4iQmT)N z$9trEujb#<JRV2|FI+P@54GJbh*a#i;x z)|-q$n3{E{7Vy+jC46!*^jfEB#eReiH@J0x5EzJv%75VZN(B&r#m+ABNP4xk^BRPD+E#?5rII&J!to07(HY<2 zqb`p1uYKHX20qv7ob2sQQk@b0uzlnnyKXt(8jj30axw`h&Q#|{cr~_|F}gLn)9_@d|eY<~9J+lP88T^4F^S8c(!!dH=)Z zf&QENmS;W#4G5;Oi7d##aDM3N){ds&R%NMUJvCmt?<}u#cp_^L{dDSD8U@*bObL;9l6L+^OyA zV3UZCmIb14zq9vVCEcIzkFkaZ!8gWD4oE|(!Zj2ahMI%>y_Vk*Q`HL`;l6~&1hhCW zcS9RoLG$D>KN7i8BJA|)?N6J=*!FGfQ0n%H;{+4qKCwgD>wF4@0UP$r>l36WdKwA=E67*j(iEZNLp=g!ukZ~ zru)jF_xJZd`*kylI{jQWgv^v#eSg1kWHViA^A5zU$g#aeP;L0v-i!8ukN;3A!L>PC zw^DtN%6+MgMOdd2`z2+3H~X^nn0jdGii%@7&&EWIF}rIM6<|@9rwhvinV-z$BesrC zPIvAr)=@Bd;l5;QR>w5a7X9UMBa-oCM+CEQUR|~S$t9F3;Wp#cW(c8y8aS^!N$#Oq zvz`0Q#mULpfbE?W=}8o@${VeHKtT=$<{v32bfyg)U-|=Ch+wh*l*oN~pxB+);ods> zQ;|No$KtPV!JueOZNz-c7;>y~X||1VTJE8-c~*9Iw&XS*;|y%RSbhcK@1SX}z37_T@hZS+474tcaYigPxz@g}rlr`V^9zT_CR9T_t z_^;Ve*=F57uCk0>G@+xGx}&3`*U9_>^wEa?#w`9COV7XCK6Lc{0;IWSJe1jesLN=a zsVp!)EFyNW5Cc_3&V@wHc7Hqr4^G5fi=>lrcvcF#(ajK*Uf}>6qMlPX@m|;k80J*&==3X z)M{wdc<dCgE?Ms<7IYRAJUuyr9-N$ZKXC&ac#rCpcu#7^YRvXh-8hQxnJ)d+ zAyFIp!qgBuM@~`gecVBhPKatE{=@;ww1F*7SXmW&>_o}e+QO!HAMx_RYl}+ZJCmg! zFU*bBpMyX!38{S|H_Io0a$|tLr zDQ)h3-H+7dz`qJ}b8}1THyNxOM7KXaTbcbZVLUA% zr3b&AWXPetvtM!cYO7_^QAMF#51Z^dW@~O~nPWc%_h^>F&xjZUbw|pEIYEsp&SJhy zRsWIacV|_~<5&&hz^14o@_25F+#Z`nB7N@%$e#l&wP&xs6ufR(kh zp6s85RLVObY*8Gu-Bb$r885Vr3&tm%fet)r+{T{`Omjx2Mh0+QFh6&~20`-Tgw=bn z=!}?BNSp@_KOU&fwO+}NODO`*U4DcQYL5hzD~1-8CzKcE)p^54QFKe=Z7&GtH*AL> zFD@Irvv~z0E;%;$k+696h%6|8b<3-Ov$57B+8FMBwSbh7-+VJL-CI(63eo=L!){D4}U)GbD&cn*LZy+RAwMeUzD**%GGLI58tE0P`V4@mK;~ z6%zF@V57II>(-$Q6&a!L=MfLL-y@~6e;1s{-|R7=Ey&5h31RSo(9DrO&4andV)@QG z@gWG!3(y3`}SO&*PCPdy>^do0jE69#tBg&OOe)8}GYa_PiMo>vdkgS7g*_a+=nZhe$XJXotpWo~zEAP^#7<2(XMKZ&T%bNj4+ zZE+LYVnM}JIZVXIy{Qx23J$6(vJBM%oN9J=Y4FYj+TxXba zz%)5P&8o>=QppjzK+@PHKOdLS=tB6Qac#;=FG8n|SE`DYGdsUXfX-^bH%e_W*F$?Z z)r4n-aKG==PwBGc!Wrvp4UgRIe0CZZxz`8b=4NGOwLHbkgdt52cVq$N-hznu##)02><{Qgy=MRH z98Lfkpj<1|4!LF#E7j{1FedV-82xp@5%thxLMKnjq>bNW0zjQ{wmR|SxZP$=wnCj3 zPaxmL^)sZ#T}9v;^3}x?ww_?*)dHKSPD1*&M!BwJT_+=TK&@K}T1PET6MD2}yNA@y zGgU{X{WxWy9igU~|KorJ>N%N{TY1X`Uh27ypk_T@5~=2i`FYDnA-HbY8R1LQOJHCA zH+~{-=sIefO^q}_A7nbUPL_#;$6wlFAZ0yG`x;etwo7r=~g^VXcV1ucM*y5Ag%4USM zSWcg=ez~93+3Q^yX(zuxr=_RcQ;;*Jy6?6~Og9gSReS(@Ur(=CjYCJ!>)nbzF{=(5 zD<$!-#p?}cr^Kz(yGZny(APc-&@l9N;6s)a#J?&dax&Ef5g#3K5za|*iHYSm+R*x` z1FVE?u7pNG6H?xM8~LiQ4eC4P3#P%%&qYs_#gD%VMQbh=J`!8ykf>{9b?gFh^i0I8 z78bAageULmvMcx+8W;%YzC3vEgjl2IP-h}8AT(hUa$Scyzh~@y&%3(x#38stC=(NH z9>GrwmrTMrtTUNFOpR#Az4e*QZyWgBv$`B0W!m=WWhnon*%W$EuBxLn#KUH9X5gKz zyVqcqlozRMWN>isbx85(&suPa-u`AAlj_adCmm20mPL%sY5$*C?*F*-_t7&M==u2r ztIA`&SUy7`hmay0MQ4W<{?n9bp4^l%X@p{5j4X3z0~8*_o15ne%gy_;&`!$a(5HV| z(u^w4IMp~C>$p3xl{ks>&dADI`>rPB#jZZ|KDO=(528$$z4ysmT;fp_+bbLjTg*#p zMJ1TtV#*hAlOej&bA)if=>WbQn>jgZ|8hg5&X~FUWNKEqunKCq!9Yn!+zL`89G}hw zqXpB$X|h^8{YG?(~Dlhk$0lTui0(C ziLO8HT(e_<+Y~-9;DPD~docGW2InV&d3@)Of{X;hSo=X|*j3FNaGRWyx~y2v>FCO7 zDr#zKreN!V*7;2>Tz6+@Z)YcGhg7s~a(cQIr^O3TxugM@S7b9GKdZDKsx$LZRobA-ix;!*{xJ(V^)MB69>z>}wK)`V@ZNs=X?-}Kn%8@AM3170K7WKyd1!4Y zckFQLq5Nf)zPn^4`6D%irYCKaTYpkCLQ~*5;RSo&wMe_>j4~bs+;Y&VXl=nQd!1Iw zdm8^LR1X*tFXOEOl_N2*Yo)3U(Jr=U}4kL9F{`Evgi{1n6(*>G~`Zl7~I5;eO!P*2?gm_TUi}HmXquKlO{y#vXpN>2W--@kFIOzIcN0n&lAtoSY1qE-BU<&y|l1K{_;=Ew_f# zwd#5mAQi4OKwvN!4sLp|b)69O9mb-b2NoMQ*iDSax97!bn6-vaGvqCm~mnbecI=-^MEd?>IA^98*2)X&a-szoQt93wPj$Zgn zrKILZLq(&^Jq{I|^O_ ziItmQs?lQ%djTj^(htdw4M)%V*B$ou!1!0^`MtdkE+t`jx5(Wu(4(>5TsVY@eTzG?yemQZc8ydo^T_I53UqZxCQ;6x7&LZ@80^%iR;3NE0+GQ z(M|4&Ma5y<;HDzh!wHhCd399gR~NN#(;yyQs_*A9bbTpgXT}S z)}!?XHjkESv?q7Jh{x*Pj*Uw=8t4(5(KL*uKcB`S{m+&(_pNX9ttF?3hZ+rn-#p*MR`J(d;jdynEm zcvuoCk!C!|GF3M3cbEbfd`)&S$$K@1^n`Jnypb_&!kBTwxjjv(i3Vr6;n7}>KhmVU zZq#j-9+2?Fo}P3dvKQa3@Hy6HBKQ1G8tj605;3RZ1|3G%3MArrQW@{usnIQOXRoWD zK9&EZ|9=+dDbe|ymbEWY4FRe{aQ979VQveTnka+~Pd8kJ9R?UM^6IcsDD@?b?0;M`6h%sg;nYt?7M++I8$p2W|FV3T@=wh`B1r_rF}XRdj%Vsw}p zis|E}kmUoe;GI=`bq#$#eI;tXGjS!Xm=yl17s{2al=B+E zgNUH!awo2;$_e;8A8XbYX|G-jJzYwzM$b6U0o|+~ppE9;5gEmZRoj&g@bv2)fkPX? zB>sGmp?bqyW3Mu0Tjx}-dq;i_wZY@F8iMHN6l^NAd>);#J^U6V1aElP{(>_i2=C4K z!<&fcr`z*@{ZINmIqAV=NC^on6FuT&t9rE1Y>+1fHR+=*EY5GhGsCq(Q)WG3XZU3K zHbUI$n6foNryxgT;90bOnb(17C=_({6Jl}_^wL#PB{f;xqCjQE_mF~6P4rzkp#E`G z8#y&~ML>YC!FP7Vs6t_&RZq~9NuMVW3W#-`mI6KQ;@Zfdf$rO+{G%z}i!d5!U72mc3~hh%(0)+R z{${}5O)sPbY;5M+Jvw|YsT{!xKbRm=PZ_ZVae5SXhZt&H*H(`iimn@HVb4Pp9E*~@ z3_wJ;MU-(&sWfhF`bnc^okTN^YDobZ$PEx`UWSyEftE$=U7zn<&DY$%_2OuIEHEF- z8CMJ75EGMI@x4~tUt0PzQNXhEugbvj~ctzYJd6%C4I61DVv3OiL zMx2peM3{+SCdKhvb07rB+xuB;bp4Kq<0L}0PfD!aoJstU#GxPpO>4$pk-`rh8RG|6ftJ-2eWgFJF!ryhYx z-OwnfHYR6eB$M6OScid~)_bUhSytNO?;HK^iD!>n!w5 zjr?^KK_0bWyB^Am$Z87C?=Q1T_u^TYa_NmP3T{-wOy=0h|96%@L+=on1)|7Af+GnJ zfiY*~Go)(TxXJhWt($!pel(g%qzc#5eGI9A0r>hKKYJ|*TBkSbc3gJ?FRBuQVJYlA zF%XXre5SxKG57A)*5_gsES=wlD-N6ihvVk_e-&yg@lV~O7vV|2&Y6E?;%^Pki#MpZ z-Ps0f5nPn{uBuEW-i?&p)!1*-V((U{qtfPrWpa;YEC-(zSg~T%yln4FHIGCZ)R0Zn zIh;m77!=HS7#gy*H4pX$8M-dCo5-u1kZJ4a*iL06uegf8F<)aRu+K~TJ*i&4rnqB3TiQ>$=4}k5>>a8#~gonSk>pZ zxqUT!pi!U8%#HJ_y-KXbncnjr)(>()4>Qya5haAP{!PK2bw8?ez2}_ozOKD?|Jv)Bx#ymH<~P5YCoLyS|9DEnepu)$ zJXfrDXh_qWF)Opx0BizU3k^iadY&j$<$0zIC3(5HHFJWzF`@q?xIY;7kOox2`o0{KifqnXk&phz`OxWKH9M*Q}y2JmQ zK1%qaGD3dg($LB+z*{4|Xs5O!YkMdk(>hb^m*Dmkr9q-YHsS)zc->{Cjf@G&YWOsq z9w(*5HDGq~8Y=jqLdepr_0sXv=KaMYX1?u$)_c|%C?#oZoHTkcSLx%mT{Izvx!7X8 z32*!Uow$&z{u&F_6X7kEGeW5e=H{#MM9jg9*o3GAUa`~A{E#aoN4$aiKa|%jsulbgt<0lXIJKQKo)&__^IMDTIffA+137C>H?&6 zK>Ks^>e%qA7QG%Orh2-yl;m?9C*!Knxc|Nvzf$GbO-+|y8fnJW+CbK~5s4sf>GAn( zz4Lrdc{$?Qd9t(If+QyRNbiIiqbmXSU|bQh!AR4oh~ycmoq6M_iD?O$QeoFAgh>I*m3;*E`N$7MAoUavaxw^eH+|Z z++06G<01l3DHRLJzjg|2`ZXQ2oQ7|(g$O#W{dSH^cAkB-XYsbTT5HUFtL1j5n7Lk3 zI%mXwYdA|_RQY>2p>dMu-~BCsv?2lQfRlDkp0E;$=0}QyZk@~tjd{%3{BkJm;Ob)y z#|nO>MR53ALmi4_Vr+r&l2j>Bqp*WH zBD8=^%C&Jmhru&|@-c{pED2XHuGroyuz&mKTfOM0^JdIOyTpE%b`LPb`{azz67=J& zn@eDG68+n}*in6c=E7^Yjf++TW&-H5$Lz^;UO9hLzyL`Jo9-RiC)OwIsRS=VnL?7* zh-Ij#BuZH`?)o$%4OjErh>KJ?Hkuqhvz#O+zxVzjmH7rluE;I8us7|BgrK{iD^=6E zJNemT$fh`|<7dIZCa* z@T*Pv3$6|MsOm1t$l^jrkf16{vx;-Yrt;V5Mafc;EN*3=49s+#o>Q{C%(?Q*&c z`Thc4l8P+7pRS9Y0(tl zbMOvyd`p|bA0!HI>4V%j*C}xpjAdS8DMpwEG&B`&_UH%UH@3Wap%#ZblcqHF9rFb9 z$rM-4@8EuI-I>a^2>L^fh(7x~E1upTrmC+Gi0(R7f;R|J~FC2;}+vKG}#d ztx9Iwh`0+D{NTUKN8`Qw%kP$5I(MR12}qGs2v;DxGcz4~hSHHe*4QvTqjt&Na4fNJ zj@%P)9Hg*M)aVVBag_5mrN$oOr35Ws@=#8~Xp4BRMzK0$Rwdt}A#XDAAf2l5Mt7t? z$QiMt(}KvK*j_Bx6*xM*w()-utj3riBhi0;e$Go3*QA*P*O#UDGUGbyj%|MZ_%n)@q5rRJUsj=fO0%6+r7J-WfM;s z#Du^9uQb6ME${fT{`Kr|{Bf%ENAhw~ z=kTV>jEOG^fPi_Ph=|)ZHb&Nz5EIRhIM~>@jMUvjvMN7bz55KciPx>#b*g_e&zfj+ z*f+(}a%DDCXeM&D(_ap9gr@o-GlQ(l8(De045BPATGD(TNch({U2t?;5$L)ve^6f3 z(bk@&Ia!iUkS-oDX{mLE=qeUZAUaK2%3BaBv6=m2>wAOAA0Yt8wb_g}HTrG3Y1(}$ z@gGY#=>Gn?i5@f~+0dxq0pZcO(K(nnZCnlc@`rP5yZnp@G2rw3jj)zMS2O{=RQ4*Y z5OUMMw)HRl*<$-Se`t#hwMLrB;un3Gs(mP#PqYa zCr`#U6KBVd_iTxat}$C+wL2l0j$3GsJov#iiO;(a>N00Ez!?GMQVKCl*g+sn6EeU=GiF9=Dn0ww&HanT{#JyWX1*_C@vd9!Zbg2ZzPua(7wRe8 zjEjq>c^Ic)HqZHw?L!gQqy*}$D4m33k0bor-}+|9v=jdXw*yA`F+k2~Hc5%LT9}ub zlW+KY#$=$z<<75$fgaANJSVtL=6m1RieqcnH1HV%`Ndgl8XpeXDeB*)QjL zrWDRD_D~FlbT8+*NlCbkmh6M6WR+1`&#n_acn#;Sa=2PS1I2-5>m;_a5 zh~`c)o*dK+uWXN_1xlAq&dv%hb$DN290-lT4jw;wQYfdOAbi!F#L{GAxT7Ba5}3=? zH#msFf3|nH+_C;EUedGbU&+&$hlUr^}*;AUdkF6rt)M#?F zE!>Ty)nc`ZGlm0sN$(a^>(6S)@T)uzBYfpa2&o`W32)-*){MK#4vsJ)pT!ohZ?4zh zgMo$-6?_Z6@geTzx55t4ndJ_dARb>_m;)(dNV@`YI12W$4{`AW1YkRb}2S?_rlKK zjU9jp-l=J-Wi=c1y4fZA#(>(aijh~EB(JTk ztUc^S@AfZln>&mueq{L00|;A6h%8>$chzzkkhGjH9S-)){y8O;;Ld*z)psy;94j$n zoP(;t#WG7o)x$67EmuwD*Z*gg>-Pr99-_tkt^Ti7w!Vo+35e;g@>1^I8bLy8P*bPj( zja0{q=QI)dXLNQoh*DNKfXeNy=+%Y}1=Geb8vJde`?$ z4r))tO&KQ|YHv`0d^EXj-u6Qwra>vUauJvXcde5A+hk^Suwc46h&wF=71F=1QaE8n zl^qlF!moalV}tulDgR!iVIt@k!_KXd7I=pasJoSh>i$0ev(T_hP9LHW{s`B`%>Mw!5PKk24Dy90ZyfXl^?)`Cd0Hs=~?pbs97QKjLwNMKJ z(a#IOTakI7(u;UZZ;w>DV!js@FAKFD-yVIsBge(TnYF%~Gh!#+2hj+yoWw_qGSX7v zpe}dzQc_^YO6SDtr9%aJ#Ek#xA3w<)5-6ge1P;f%l2ows*LV@Dz7fK|AJv0}8l~_p ztuFs%lCg)}+N<;w57J+&pfjv;wlYD7Y@YTp^Jg?xe#Lv*>TWvWs=fN+0{)@6CiDlT zI*sQ^0QjYe7T8{A37fnO`Tcueu1Pr;hp|PgKHIVLH+DzUQB&Hy7j33-SP1J$_erDV zyJD*!WfUTw)fe5f3^cDCtJ`5ir0HqSm(W~!z8Xl4WmRP*FHS7-t5biNz2*-wj;G4anz=Y8cLHJHI84Mp1w(O0wRU}gg7zxEYF|-W!IcO88_d>E>M{9Ct&rH5 zMwFAqDN(Cf>F5-hRfshf|Q0B(=*^L$Cp z0Rqu=<}S_G#b@1UGAGF$l+g*_80W*L2Gc~wn3ntFe*vR)G`uhL_^~6Kt|n>gx$+cK z+d?}XSGy0UtZU%5LbcI0wsvoYJbZVm8S>GPvbicQFE&(H4U+ze)}ur(v84dzE~4Uz zrq9@o+GXS~$RL)pw$XBES0N#AVv1VJkSrOAmRa;k{6x$mz#JSCRe1+`#)-YjuV3#T zUSglZEDi4BGgC>59E6a>#u3P~tJYlKW+SPJ3A4eYAYfc^hsaTv zpBP(QvDA+=u3tI_+l-}|_3AsG+ao!Gxn4)LF)=Y81KX1hfUn42OSl?%#qVcc509V{z<(j<+?NU|)_w8S&CG#_T74)}p zjE|3t--}f2G#=2>KQ5|ikPQ0%qN^k2+wgDT2qvhjKajf^C1Gva%9gViVG`rxGwX>W zBci{9NautikSkqPt=dQX1^F31bgSc4o`JW+pF?r|sx(@F)JO_GlXk+@L;bJ)bAE6S z`v-y?9DcUGu4D^IdVl;$uAu2K!Ty&Tap@Xr{_FVj!!GK-y3fBq>H()l=XEYt0F?)1+4X5@A*ATVpF^;9Kg!m7rzv9-IqJ7{z?$M}bpWVaH- z$@z5y6(ya76P2#DwLI!(8q<6<1bLK&-pwvu`Kkfl;5sr=i27>khK73`1rK}zuC`6 zieDLb)ZX+r$NzeSaG}jYzeM+9s)BswRXRDmeGuJ9dK`as`uOCk@;4FqVIQP4u>|?k zN`08i&qrf93hf&(ckz%u(}V72?XIl2w3LN`RrI}%=G5fEhQoF04fo4%bqWz-!A=`q zSjg`@kip$ekJveY_5B!4Iw?DyWn6?mCQVplQ#~#&cwPP!EBdiU^((fj8Jk+5sHSiS z#C4-hN5b#%zKG2=(6ESdR2+`TBql~FDLvOHR7;T|_m=fLj-Y&;fF*kT`}5dAX;D#} zZl)q=d~R;ea0~CB@{2$Jv+JJ*jq0Na^_VoHGer5gtomjwRd}R3DbjFSl&RUH1Kq8o z#%*sn%jqV@u6-oDd}{W%8JcIz0F?cmhV1bVQ1=Is&6|jAhrvO`P=G~j3Jl!C_8EbR zn_xVC{PPj&>DBDCvgnl>Ct=U{PL98$F~a08OL2vj7xmgo(1kF6av+^V>|zWl>ol6} zFQ1>+#x8i8`P$hH9K6&cqgHP`fj2(Cy1pI~MtXX9)JtwSvZR6Qfw|{@FY(}?W&|j0 zJ&(|csY8!i`zOxmU~r-f5kC1cUx^6Yu@^brtCR*3r6MO3TKIv2fSX+_V6~5DsIT#* ztI8hyl$|73dX7bs!ldcDSqga$6lqSU7n;2=v~{#;9(oz#n&IOZ$L-pAr3!FUrG?Dq zS|t-Di&h}i^T)RDuGa}V39Wbb-R_NfvEF&xxlmgl{vdzAw#Tc2_U~r|YABBPHIz&C zZbES5ggy2M5pKJ0A;!dsK7u)bLxNTvvZJOB4%Uem@$qvhX4>_y>0-!OIFR4tF=Z{> zji~UXdweN;?%1)R`>W zy0%~~=39Bf5#i)Ye(=m<6ynk6zB0Wgb0c@`_$Pn4b}z3yf*>8-t5d-x&h**mDRRBL z3OQ$?mRiFVhJ_g24phDB{jOftHw`t`7q91W-Sf)I%Iq{Jp+_n_>9ijq1{R>Nqax{j zq`aq}$Im)sqMlviH|D4C`7L~gEck3v8ka1T#-Rr?+{)W$|FW@R_t2c+Fv^Z2>o04F0Z?;&ATa1q-ZSAn}972_ zH*xg!xn+>r+GqCm_KCXfLhsUD&S#9L%zr>WFn{v_Om?qedQiyi;3l> zAYWZZiuTUDW4f|7{P_sNIFE4}IUQa$(L2%C-$zfz$p1N1AE~RTsA%hn2-Wm=`_j0A5Z`^B9B#zuD$RDcs zFd)FzjOc?>OqC=XvJ)c(1$BVZ{;F%^uUxB$qmGR8e;P8%Cm?VCg!S}i(o07puxVde z)wkA>veBJWr;EiTyrx8KkLox_FD-EQ&aYWKW%!wC*<{{}!rhIZuX`D^m%OO*GDI>8 z)p9va4X-lqXfy+TS6aN_^F0D@(VC5*txX|M7<-?j$vq_#vE;`NRjE;5^wZ3B5DNPk zQE|^1s23AsTPZZpoDpUAoStMYS;WUg>u4vcGBa(uUbq%$JcpBkN4f}D_-le`^+e7XXdg8f|z}6S&Sf;|pbS)2bL~9F0=;kb*8!-ro(w=nw2H-1dFV zRJWP$Maa%+P6#i1SE#*oh#EPH6kq%p@I5Te7alu0keU5sr?IkXjX!DO_zYAtKT(OXs(q8bEX7-j$5n}Sa?#%eQ&Lhp+uFp{0d$j{vaa_xfT+-;raix`bs;>pJl>?!jYdJ2ZHImg zQo?t;bucOT?UAXRw_~r}hebi1%@x1{0!d0MKYLe4wk7#Qb3R2a>FK(GkEhT2^`xy1Z9W(0N@EieNKOAd5hC83@f~bERWcKwn8+tt)UCEN+F5M3PoR3G6s9B~t|%bxtsS z4f|Ks>yr3BFZBLAVbk&r_pMT}QGmlr$@-lFeCrmaxq%sgU*_g71f}#(-K8xd2P6u_ zVPM)oQ8a7uFQd)kiziqlg5!8ZReELmI;vc^{X?xW=?mz64Ys|dT1xq=Y3f1rX6%3jX~8lPx!DhiZDKHdTQ!d99_a&Ju&_F z9C;NL6>qnFy}N=kr-Ehm%sYWH)1?S7 zfAw-^+h_QCXLonk&u`x4mIl}9IgyZUxQn}P+Amudt+^&9W=`6984BN@$|Gjp_uO51 zUKVJDSY3!{?vMdu5xi@S+qsMn2~89oiNWqD3xjI>_=ytPmk zKLgIB4TZj8BASb;S#6N*HMjNs9Sf+;Asp_`XRzzdjeK1XD%;=i2~?82g=#0qG4cEE zIY_1Y`$}2YqIK&2wD3ch`Ao|uY3uv95M{Du~#kpl*D01Ou zDs@Kgf-5H1e){L)$O1XF`f8`B`u`)&D8~7<10dXn;B*j%L+Qw=vTDY8!bEO)6n_gJJ zFS_2+c%bNabx{6;0XDx>Wqcjgv>tuK;2%OGStnTKa?vS7kvReLsY_RqQNM*tCZJdD zT48|o2G#M4&hyY%j3Y1@@#(2lwj=CnB-m0qvOShUP;JE(0w#E!K}k@PiVSZhg`EJ$ z)Vr3Z46vODlFaY!y4crk5v*)~_s+ZfA=Uz0ePvPzD53AnKq~cLMw0pEkCB*Ge4xXK zU`=4=)W}5iB*r%&N^WK+EvRq`<3mf+#zMU`2#Ds>R!?#;Kr%Yvuo>PrMK5_prM6wa zXI6H1_DBz=rRq0Wx}?k{A83@xDw4-nn;s4FDQo_$k)SoVhh|y7WO=%rmil8y;YE>3 zCc=#dP}73DpD7%E^N?vlOh2}Y`g?P|BOdk6=xJ-w+`O|pk}vR{6Vbg|?T_^TsY_^S-aGES=7K(iXC zK>xiCEy9=oSJhgR+WG$;zzUI*8@*k8-6LZlBe+-uHLlz@;9=12He^Y=?s~W z4&E@-y>sM!N4oZ2pnNuItw3?0+Ek_k;f0P&(@B(i#ZYsHs+TNh?noKn{{iko5h1>0#CpjEMv~U`C!Tbyg$rSta+eP#Aa|j zdctO-@7c58ONMrZ({1if{oBn((+t|9nLnPmzmO(3k990Sge#Z)9mZ0CtSxexRH%Wb z6i%k_J+W%FRT$nkRCrBCm7kBSpD#O1)b}(D&c6UUW$kq0ST1;eh3Lz~{3`)%!o2U5 zL5mTOXLzqxfgF!yUn~TOIL2iNNu^_aNM5?Pv}7?@RWH1_#~=T)v~q&No)xh{!S%x` zu@yPSp=zP_V^pNQtzFNqXV!}i2GcMQs*lgvj}|?5KR@GJh{={Uen`5Rw&=5&PU83B zOfongi-qGzM}@+V>kHw@mKO^)Hh~n!>eV&hOIvM`@P-zpL@o&*aocBwlb6Bh1!w*h zooYeQgB3eSRn35DIWridAutyeM9JEKQxJZ)5~Y#Dz`#=cxWl>SE)jS$LmH>893l2L zO`oD&tkE9=A3Au z8li%`)|jfP>pqn(57cu}3eu@;`z(ig^_@$hbTk7AQO4mxTlH6p+5vd&tE%M4u}n_> zTg=GkWJ0bb-aDPWUP2hhS~Imv~K9y%%w}D;-*M`#ahnP#bbJ3YA8(QHM%d;+%DyIW1(oH1vsTT209@j9B9SM@HIwND zx2Ll+xq}Hkd9I#GPrjh({CWvKePnnqaMCR0l}`t5fE*vlo`h$NdoS(C7rRhqBIwgI zfc>6+w}iat0VuOn-=#qZt#=FOtN8fwCVe%aJt--Yf8xHMBpw1qinh(pZvF48(#>Dy z?5b2>Ft5=u;es$$1Cj~(xS-5bn`Yg0$w?onp)8%h2L$|pNvmO)KJ%?QDPLP@zKH*00Yj@~5-p(6~R7>tJ6@0=PxWe0M&`FOs1oiuaiaG}( zJ+e5s^*Aqes~l%bqX*!DwI4q@RrB(dG;f9(+j;eGL7_#c>Q}IXW*sge9v-dG1w>;W zut>@p?UbPz%~gV;y?WFHx>dACMKI4(uQ7eM?2Yi$p>nX$thB3BwURwig|c;cszii~ z%Om^f(6UKm7SmL#x-C*B^-O+8qm70$62YtbW7~;|i5f>o>hx@#XR$?GE0C#4zUC}v zgL}_4%<};l*!`dNkPI+~;c}2->u8%5%vPazBpzn?j;kgz-h!Hn#C%SgW;z5c;>hVB z%}Zs68#2Qm8+m&t<{7Qe8pa@1&jj%}3SQ;$_xCp&*%y`ry$p4lPoS_boFqmXD>r}m z)@ z@F!@z=GNd$?B~6wjW!0Cz?^RX|ms_d}03hYKZyV_`}`hanuzZM_b1FW6)if4t(M#?b>QP{wxN z`8_ZPom!(LV^n|u(14=b+Ky`sUmHQP4SRi~46qUqlGqXmLIUZtvc}JkpaBd7=d4=iOVi;IlP9i{N9HGUW_b@`% zDkart8D{Nz{pgF*!4?s|7X-d|IG`szGD1F<>-MzBcm%_aAtcGohmk{~0zan3=o)!a znLg-Zj#}{S+o$<((LT~e0ule{rL+U@{fdZVvh6}=&at^(vXPMyA_|}L?E-+V^qhVi z(|O%=Jpr=d1b(k?x1hH3YbcU~#v2AG=It;sk~VO3z5;#(l+og2*MX6Mc z=&yKB1grg$ou={F2+wmWavJqaVbi{D9_vZWB9|qv^F2}T$+T*AXuU~DEOxV&&Z2wj zsL#DgiC_Gg2>eGy&zu1>d*+ZIbSRrUS}bM8Z_HBq+qSbV1D8v-Bo~vCN^kM@HR4NU zHgZw=+yz&~hIfo~WzW1Yg@b(!2J~!tz17U1Zb*ZY*^4H)xX(|mh9p%mT@*Ms%uSP! zTR%;4m>h?n3MV;2U^Q1!MEf9TX-bpnQ}F1J=j8D>dA@zxebf_8VgJ7RPms06`(I|j$L~Np%y+D}N*$}Z$`M2jmEXSMVj6x( z+-HG>pZf`wmZ!!TRi(7jWcUBD+7!Kl7RCHvIB7g-`i5(ouzTm^_10Kxwl^*W$wmS^M)Tm@9lb_G0UJz2*@*mJ+Rz9)eko{ezOgyHU5m${ z&#!1a-Q5<2KB$@BdB4e{kyc3IQ5i97%fEUg1IoA?N_eGJE*$6UJv6CB!$$4JgxP z8Op&o#~1&EvApzodEa6@8_62A)e%pHf~V-Q8ZVqWq^B(iAzL;HYUG>a2LCQ3oMZoQCFPu4ju|KLL1 z1-(}8Hl?5j(V;Ci_B!|a!D6=>Ho&~V zc&>bAWt?_Wa-m;9mF-?wImyeZGd#N{X@o>kCFyW}5UK;S2CM$q*|)W&R#C|_*@q^cT%?%Pr{3Q_E7EKgO)j4qktDnk#02_ znxqC7iDG!Xk`Tia8?ozbwmrK@Q3 z3n^uGeg_Dwesg{k^408?KQ`pr>+rGLyLU?q{pC$X>nrlCpVg1jR9aN*2X%% z&mFR;lEUvpE-M2okvymyE-I7DwGnM_93$c_mokS?i>^zLI!_x~gw9@^H5%Hp3I&81^h&omWPc@UXuPq(GN#U$Fmw|x9&Uf*gWy~APM>U#?MVB~ zaCUa)B6s$uVMQdEr}5931$tJ7Uk#v#T^c-A{u5>a?w#cc?~1N++B33}+?uof=)b&) z=o|9f?;x3{E)@Bd+KQSSMwxFXvhxP6=+2fk2|TaUdoB)yW3C|A>HCT`s|}8+W26E` zV_Q0zcxsYHX4oU29+#g6)#CQZxv{ni%N6lnUG($P>($oj;C7)18Npnje`r-7IfZwser-YlNO(AIbXz4&rT%C=$UCx zR3x4N*6FE-LMdx#>&ixD>?T`EcN4!N)YRT`=6sS%E10)Lff6#kj4Bz)z}fMFd$KVe z#S7OBmKs$i)F-8xHW&N3vvV*xIW_;Rd+O7u@7?w8n5Or&tv0$0@gqLvLQAVY^&!W`|Q`6k&_&6S8*vbpGQ)#PaL%$nSd8^<|wvTEVVI%&%_or6I= zZ1QDlAtY*Nipf>EXTEejQf6!-!u{u5a%@FKYOO7d!}hz<7PxU9?s2kV2c50R$zdo| z>Yn*E8x*7-#sd@sQBd>)0o3R}>psH3q^dAx3wI8KX9fJGs;3OXHh7wTiJRR-N-0c6-`C!HY5XJ~3_WX}G!d)1u!q>#FwVJ2%YJ^cv)XD< zaS<`mh$2_OHXv4gk>0Ao`;|QTO$$L8L8L^MRxckj*XM}TowHA1JwxB+9T|JepM2x4 z*6M7=Zz}GF!8QuRADZS>?VV~SNG(dIUrK_|>bW|Cr0ii$rq-32zzQfm1JF+$Sj6_& zIo`{cO1n!*zxf+$YqOSByY8o+qM($&`ESbyn&9Uvw2Co$xFe3AmP>!q@@c3S{nP?Z zY{uO&K@%_4Q@ZR2tyY+yLilUA3B$v9tIu4h>1UWX2tr4ERB9OHZ>!D5n~I^TMFb=3 z&hXnICatSQKhmE1>C%Zr#?ltb<2fm*Sb{n3ghKAgg`JAl=_c;8$8NPiYf}TQu;*^B zeVN>Lcm|46<&?T6{O)kQDNHYQRTi(2+pw(FsMG2%oNl;s3p^uhqabOy`6z1NUIfB% z(TJ5FoBu$zyU(Te5+~i`63@t{?Z%^bM7N=N9|T3E8vUImO*-+jKjP@a6pqLliR zQQeknqOD${&mrutE+slK$`JI>mIuT9}CYfS{yxdU+saS_N&uh|_eGr^tbZgq4?9FwQN!rAr> znlAHv$<%xS8i^~j!EHVW>LHVvnntLjp*eL}X%zVe@TUQCJ_Fp&p!cUUiQClqnF%GM z7G5s4ZhA*qD+&E&I8riR3ItK%BZd+%Yhs<76zx4dy(1V^!C5_{v25|A7^gyRO~lqo zp-)<)ZWMyr2rQMGQ>xvS&Vrp5KBCtwWa=rm0Zfw=g&I{Pc~uRP~{~&C@&rkk96Vl<6I}B7S9#1EsXLmXz-rAHggi~On zpm^&*Wk78&V$1CP@P;+{j}F2w%IkiXU-L68-Ay+hZ1MLMV=i5!&lJ_}G?%1IXjGc^ zH>bk0jl@rJ*l8c&G*~`~A(wmX|Kb^{>{jH;zB;)-qXi!U1DY6q{8hD%)M zC8d?KGYJ`s0;JDRpip7!CM2}kKPVxsZjPYF`)G`U^l^kT)YEZ(ne@%Jqn% zME?@77i+j1u87R!=1triKC9uSb9(D!POr#XG?CyNdsJhHn|EB7jIoL|UcsGnIX>^} zSTDA39Zn=2Wzc)COhuIQ26B5lx$^V$bHUfyyz_=^uMCy9&qcqA6^qgpFiIsT7Z+f- zK1M8|By_*%1^Ezi{e`cI+Y3`y#}Z7ii#KY~@bxAh3l59x=u3 zV2)|)%|y9boPwK*+x6*~^7-xFto04c%l>3(_6zJBc^Is5vW}MUe)l5y@%|6+_ym#U zplExdL{4%OA0Q5O?o7JvtZ1jlN8M?;<8vob8q(a;L4tFQUtm81I3n0@UG^>)^D~DL z=P} zgKbk zOUpE`E+N-JQfH|7DX(dKyv2v>NOQW`D2EfS(2Gis9lC-c23z2~H=?E5E^DDMwu4kr zQJ(A0{a~iuDE;f-7}*nW_J-GSouND!0n^<{4La{ES!fDqjHt}BQ?PMEIcg3&pt?JW zAtC$_@J0+_qD{cs)|5jpTyhTkcRjU9-6OAu*u=yQ*6WSMH%?%B#?jEz>58HTdA8;s zD{Z$(uB`z-F`XoFDO+2;)VbI{1cy255jAfA5dC+8!Ts!r%Ivj3^siC(&gOo>1C2yc z2*4UIWILVH-Xd^&(C(@!V|GK5GD4_AWcfQR+rAJU86+n9UV7z%0j;AWfS1A`Rqbq7 zCfrJfm=QC&Zsx?VSp`VZ#XOCzdxh?vRi7sZkBh{F78x(c53wk}u6w13i-*6y6G7cI zIx$;`T#$N2qbeL%zgC2)YTtD5HOdENL7@`H{4J{X{I=|zRo+Y0-rmbvES?|Y?8aZ$ z(HAiYN^_Bu;Wn@K(%%&c+}(Uxop<~d6cn^!u;AK% z%hqyf^b~u0s9KDAwhwgORa)9SVHX0_oc~_zA75`gJQS~dTdW2rv)%OHf=Xx?9st zVp&ulHExJ3*Jnvt#`ML75TbA)V`#Im9xkD8UPWS9(KTY;zHiaPsbOw*rL{%<<+qd^ z*+fPZei5(CnVXFS@DADJt!jhSt@sPARaG~=8GbHvU8s1JmeUW!aNymCiT6`;-(-w$ zYNi(#?H7-|4i1{m8+N;|5eNW+yIk~vtRDLGC-T`XR`^uv%IIiCh`~lizyzCfJ3GRE zOAY+$_+ljb>HW@HfOs4Y89@}mbuT_Dj)^cB(n3pEE4i41%`U;VIkABU@$;H;L;K59DKd^t z{p8#ykDt3)$`VYUTt<>KM8idLugQ%|_Ghu=1Daf%+1o5GZWS?Fw@ZrSSUPE`j>oLe zOTG6Jvz8@-B*gr9Ag(+ub!PMJcee8c(=%hK5|T_d8hG?O6b`X-L{Yc|B#?`u;__Nn z9al*W2M34YchN3~OYOh_q+6{scwfS@S(SkC3qIS8PW%0>Eq#G;@Sni#|I`q{4Y;UJ zs`M{W<6iHm;f&p}^Iih#=uV0PzrDLN4rBn@7=B}f5%A-3GRZH9eRqy(Yx^X|o{hdX z3BBqlMyKMWf1YR2p2NS#Ydu46e=vIj4@1#Wo^pT(OX9w0qBqir74U5D(z0r8`*^>OItgU1J`{Jz3_?E&y(NdC z3Gb{^JC$p7A+rLwqg_|3Tq|!d9~c5y#gJL> zT8Qhi2PtuBE#oWE77jN1O4Hk>D_^adnO8vnQxL%Zr^pvioBw~&FheWxVgFUH6Hr@pT7Ql*O7)@p?m}(KfEY-zx&&8;*Ov9 z@<)KUk#0}uwjUvh_Zx|i?AVqvoH>0)JYN%hMtR~k`Ek3GURXRXpjpYsV|dKyvc7)c zYwqk<6gyI9zv(}J=TVtkU&0;C&-ZNg-ShUIa&CSfIpgsgd8%g&0yCXSOO@Xrn(W)y z&ewXsXYa+^Fk!BDLROO;Tr@<5q2%ZpnCrlVGzVWn>un@{Mk|4rT9p@Z$(D*qpy02n z=^G1CBIw18kp1GVS+5@{WX1MR;oDWjp}*I2EKFq?-*DHL;{2~B)KC86926+-O?5}>tJ*C#~r*zK`v zK7Vpq!1NgWf#q9HW-FxJd9e76Jy?<#JG4WP5;w)0I`aCLISV#*K3gxg&rQ(0qBZ5t zWH}=5rj*{K$C=gRc)ufSwX&z48Xe`?Tt1d=RZEsh@ZL#L+vo1+J6D1)Ur?X9a9_sA zrpccTl3`9nmX`8F+HUkUyp&f2lh|9p-?^P8xjaXI@-kJuhZe`?c`z0|t1QliX~Y>u z#0z0m$o9YYQ89+YmVk(Ho9e5=GO$U-;OVIEWjm$(Cs35RGrmS(EEm zI4+K)%s5#G6clmb>*;7B+KVgVI^ccJeOcRDEe^IsHMjr8v`2Xecp?-+ToqLUtNLiD z_;Z<0iZ7lq3Sl{g2#YT%yk^sDmfQB0kO(plJ<(a4{3eW)$bBU`v$Y{>Hvv;wth-3x zjQ(bSNh4ROjicuv+t|nV(NFfgr1odw<`?2_MLFYax&46E45Nd1KD#YOn9%0hendsB zwZ;LaNx)JeSuslY4`_q&#KEGo1Zr=ze_}?%Zdlw*{r2j7dVhTKHM5BM{7ujEM#-=Q z9SXMvsonhf)1_(=rp!eR#BK4&^J3bWd7^ap`sSQn&s@oWot0;Jz*#}T>n6m-M7txF z#`8E%kv?cFzZ|r${Q5ihYqyAwyj$~?f{wLTXz$G9f%esAtHtkux2xTd;Ps4s0%2c` zhJuMk6Yq_0u#JSbwYTH(kZY+p}d9_Qj{e@Yla z0>!-)DPAa&;!vQ_;8wi2Lvbg?3oRu;fZ|SC+$rwv778iu8r*{Gcj$ZH_22tvt*ph# zf^g34*|TS!d1glcoXLWGym^1&IHQ;-pUqp$9ul=pXR056b!k(e{NW>u6QZ9V7+ZI_ zV_8JS8(X{m8KP&-buxW&f(RlJ7W?0~-9q~d2v^=2i;Lf&gUd@+#=3yI(DoS0U9qgK z>5napK^g9s7Y{DQu9T&bi?ho;F9YlM_uW-cYHMWhVd5l693O+1MGYTULzsLaFj$n_ z5(au;83P74A#hE$>>Z{(haQOoF`GgEN9WZGEB?8NbXp8NlRfTa!+Nv)uL5G`4MM;P zQGlZPaDKT-Ow+iVM|!;(Y!_CPsOh6#|&fAG>Qm49Had zo{IEYJYJqj?FI@3h)r`ixH^^e-=)Z_wx>DZQGG;vdbQ6vEHea&SR@2#30`<{3*q_2 z4||*m3(0{Qu$ru!7p$9WD$-hxdUe`s-ipolfP*pHm)jetzA5x z|D8)yKVW^0<1a`y&%=gVN|q^`wmW8=mSg4f9T}@|h|jLi5|t>^S+b&qOOtPDT2)rI zVlF2w=cF-`Tf*ZPXei<(z^_&feJa`~2f>K^Aa~{e`ilQV65|i-K#WHKch~>%ZI-jm zMazQ7#ln8@E@CB^oakFFCih4>(~FW;tC?za%grb!Ki~E&>DKQPxT+aH1MaJ~u}Q+= zxk-tgtYp6lYtmz1f6>UUx-Tm!hb>34k-nITrPjAUjp)+XcHp|J=(4Ime{IKC?qkFY z?jIxvL^GmGI7B$di;Nq7fP{E_HSMSB6zg#ZqqQ|1&QW`~Y3O8>5(R}am?2kvGhv2e z5>$I@DBlaHVRdERuY#(w&Jm>f&X#Stv+R;}YKh$Ha>gMJ=J=&_XV4Z-YbzLNwI?Pv zY~WQfq;0?<$NRM@`hRNy29^Y1y1xWGdpz#}W)K#2pE-{|spAg=dXOh25Jd8VzCv@W zvzivyZ$;m#mRc44(btVc2;#Xf=6O^3h-B}EsBRXsWKD6-zio;wi)jm9{ZWE;|9~?HZvsaM3v8yn*9#KuKoCay9dCtfonbM$HZ`%IQ z%dbQ`3+*J7b_G_zmjo8Qhp7G>M^j5Nw*5T+dFp?$m4CP2QXsIU1ke@iX@7j!HtPkN z{$**~<1A-+OJvZMbJluj-f$OK>fOv0x0L&x9=2xI%JY4?ejpF+;iF?-eLG>ii&)dc ziq#EI7OM6GBd{^vk9pTrzywaCb?dzspH!KJqjR2cun?@E~3fX*)bY1CE zng>MHfcwt|6XN6l{sIl^DPn}Xy_DLH~kU+XoozmF05c^c6?aK31y>3LOE5Sfoa6Ak| zVMe-w`-=i8c1^%)5zeM|C!eqLh^i^}PB*g)zBygt_40Gjfn1lG-l!ZL?Dm*!?DF;4 zkwHy6{tw0fGu2(x%mOyMlHh1IQ`Y^go$ZB$QKzUG`pD1Q7ic7T96ks7GSa>heg(XH zWvU4XA}cfUN-JGmB)5u7*w>~q?zy{2r<5PmzNh_oIRMPxfUBN=SC+-^oRU$STigzH zcov38w3aZY`Eu#X-xFbqvwv#ENJosu8BrwTfL>LtKxYVruO{nbql+8+**m>)>v{F; zg_*x2{!bQu!R^D{Vq@?(8<&AdI>@2Vmy<>DYhH+ojN>P91mr2r4?q=ia>5&RZ+Ad zl?2)z=F6OQAOTj04j;j<-*`q31q(G*?Car5rf!h9VOn8E4#f;ownDz{SAjo9=k4&B z^}WWSmcNgWj|6-SEP$}?e{b}?yOR1#oL8wM}=+mERo7>{N*(e3D zf4fK=6DyWYsPhwBp$DcW?H&S~G+VlhsY5 z?QeOwQp1H-@4evz%yuk1j53BG{VJn4%z$X7JV&Np5CQKK)-9!1Yg_pR;f~$sUYj;a z&>m&`Fc;Y5qyKYz^en)XM5EVB_elhRP>;Ro%^+Y9X&kKf4Cnvar@rH0SrMiJQh%-! zG3Y~Z*xYX4vYflFMCV#1_FVl&z!^VBVwKDI)=+Wvv3s(<2kw_rrR6P$=;^uH>5@BX znOM7rY`X`IK3@$-L;h-13hNRH3jG(L`FA5$^gPlM4n7`HGwv#Tx=TdQ%LB=`$d`U4R#tm0?Luw@PmkqukbZl0 zuBJ%;)U0+W7REdpGt9BaZf|n~+a0(OI#JdR*_RBbm8Xf)4Yaix7oXteyc}}>ztt-F z3IbAI_N4z}&KQpBH~}x?oI9gbpeUp{KjW~Nq6|TKcu^MQ*R1@Y3QXp*leHXakhgT| zUHl|ix8=A~HIy{XdMYKW8EShlZ{@T*?OkJb(tT+E#dyZFU7mvFRP=QSi76u4bz2N2O59CXf(lH_be<##3M|L=v{ z7rHZh3l;%rg8G=Mv%^+0pD=Ne#*@9xq>K9ksJvZK2)atSAf-rhZGe}M03=S_u8Qtb ztXXPJ#HRLUwz?=9(v(z8s8&+9!NXlmscJGn?(*M)E{{VGFPgoNP?Q6M$2jj{cw82p z%4w*=t-RkB?U3~^f)7(?EegtcnewHt(tqt{KewdPuAiQo`taia@GlLX+#MS;Fi;23>Q!u+%Q=HF^S-{5~q8B90)oYZb$v-q) zy(Bb;cg*bbL^Nv;8V|F5lIf}VLT!Vj09vX4s4(9^Uem5^@=FdYWT7rK=Ch$BGXO}$ zAh(g|N{k9FrCyrZTdCw9vsW>6YE@t~@u-ZV@4i^A;uim;X;ooN@8j!Y^VK#-RSBqy zUvMs^kbe)ZG@=dCP)^S^?chn*KhzK+YjKoiZPhG2y*`#h(bCBw*!Vt_xRs3v>IGs; zxOi+_dc`2_^}VrJQzn*uswosykTB{q&$E3#Q0N+tlvUIxU-i!HTXqlSf7Yg?Z? z6Z^s95ur2b6M~HZnX)YzBhQBBufr0YpZH?pid{Vqs4%{;F``H)w>4j_bdRecD$~Du z#~>cI-2prqbh&hmNx6|8Ki<})-Qpku(NMpcM>*;zL z)Iz0)Uzaf3S}TWqiAgH|jkV6IJDh|stK38A&c4?~0$kE>I*dKxfnjMp1!>R^1E`10 zb6VqOlUC(|Ypi4e=%%hD;#$UA8B-F1UQ%>w4?ho+Xu*2?fwp9G=5F0#NXW6nC@;-i z>;R+J3T;%ba6!MNkN5m&CeG72?(4ND9!uM3=4VSM7^dO0`Rd;%IvP_)SGko8Y^-KobX4z~AzyF&5tRdUQUP_Pb^9|HyQE!y60j z961*d4ve zLL)u$m%$X0QpHQP3QIt*pu)C`3-c4s8x~k)Gm^?&OxNxM1!%Dg535~XU4Lnn=kx0B=_EJO91D#W{ z+m8qeYOx0+!>0B&8|%Gsy4XVrf7ZgE0THU4Z=goa^3}q1R2NyMQ!oORGH3!Wf0(`; zcnl-BZ3=yp_7Bf!klyTA;tCq|yoa*o0Evdmh=fC%{2YQH?tsO6KoA|jm$5-I#Mp%* z%9sSlw9DWN5)c@u85CV4dGP?gCO~njTTi4IIL~F0T(feP ze%O*}mm>E1Rf-nD-4WqP-<9`n*9DhI-OZSy)@oOI<=peq)+=ecur!6`{l@DeFaA;d zNEHC1wA4KO5YbOgMEAy00oIcuesYtB$S`bdY_dS1N@pyn_CZBHjoIu&ZxyBuu8^$_ z@Y&for`9Z_%HJpZYQ5-qg_*-``q8L%UpcbDa1L}wj)7tQBBFpJAp^|kfIiZN$KprY z=KPI$bbgQC=cef-PLqzBuDP#Y`r>M#F-q|~n$d)ep!7anoYc3HpLfAlgZ9ifMo_wb z%{P$dyA%t;NEsL#6+ohGDagzrJJVIKh>>9#`QgZcsaN94L|*xZT;(yy{i}?oS9;#3 z2g?;n53S5SdV|8`f*zm#Buo)M7p!9kXGD?tSc?==1YCMa-3goQlf7! zOw^&H|Lzlv4e@K_oMMEOBRRRsk)ra&R1<0+Rg&bTYmEk=?H|+AmzT)qYt)`qHvs)| zEyo4vMv)N@xJz<;8GeVa!HNtV*V@&um2hV`Nru}crbUvMF9&D~dh8C2PCM|9^Fny4 zU0vPytK3g`E8KN#-Jb+Lx+svie_eS%ApSjb0{S^tPQ9Kxo_lufruf!=!Y)eJG~iDP z*Y%6ZaG&D8Dr&>;z-67@FaW3UWfd%$Xt}hDcU{WLzOuBxKl-QZ4O=_Q3Z8uR|7CmK9$Fy!sOK&njsT$XMIz zMHdYH*e|_a)M#06N))}%K*T?Er(bTr&^H}WU%kz*lF=6JD<)0Kk zKW={D0d{%_APPpMM``liy7F*3yMt9<`|#yqft;I*wYj}QzFn2TPrzKX$pE1+LnA7J z+RTJK@GqEb$3zx^>iRTKxqQ5JcpYv^fmdbJu+kkO^w z@9ZXu;Sl3B$(g;Qd%Z`GYvUI(Pz&saT_hp|hU4&2YopYe8gLS>^Ledxke@y+sk%UU;a(p`qbP zwNvbLmg&-=*eNz#cR@o*5Y3DDe^$CF6>3kPelOzpzdXv2*&7lYPoUv6D3m$|+T|lA zb{Bp<_6KzlmSd5+q+dQt-$>BLpq4O8L9D@2?QC&blm|G@gpXZ|a&yn>3^$BSRjfDL zFzS`Axks#b>CmD$Umm;JtX6GZf0aIJ|HeeiV{&DFDlNlj z63pWMGzfWeq+NdUqBvb7s9-&kY>t~59B6lb-idu=sfb8wnc>mAJXYp&n1obxY5%G; zUn}tyMo2MPU+B12+S#B~e6xB^ zT4#{Gqt_7lQ>CY{2@@r6r}Vsim!+fR5TAl-R;s?yyBKfbF$@Iyv(PMKKZqawfi?dx zayGS>=vXIHM3Ma(1lDtZm_5Q!|Ng`++R8Ob7{Qs#E82@P;AY-PDafizHMAFLI5E5j z=xEnF`)Q;)_i?)}w%zb@j$Hy|e%`s#_vSMZuBj$~4RD+L-qg_>mL#HV;hW#bHah#^ zk?HOCt4R^5X1Z_r`FY>=;4j+x6%5)7s%se5`WSljhhLMNv}+uH4^Nv@r|WwWeEAOC zq;WjRy?c2noT(B@a*jR>zxLkH7JTgI^XF-5gy^*DA;9@5PF*jw=6YkZq&fe4U69`G zhR97_OCWT3WNOMYsn~}rsdHuZ^3z4&v;Za^NNL^kZuk)w7K0^e8aw6@PuPql8(?8n zI45JoCKELF8JLs8aCH|Xm&k{Ka=`C~Oy(q4b>8zqd#qHSDog+~1hN^AO<4yf>Rt#x z{qjTzu))y@`W12}OFkJ^R~-LrGcDI22PP)~j*%DLlJIt=OFvQC>MEw5<~dxz6^b&8 z-&&GP*Qqj|hZ@5$LBKGC=iD1FtY1&o+_rSH)ijn^(we%Cx_LFFqQ_Q|x(brxc5v6h z9hXGfo;l_(!}cdsektQRPL3{KtBXjtdnB7l!r7wLn|?<#`L_h7f3wtVp1bOVY2!n6 zZQT`{N7wd#3rh$@(hy|%+(zhX+l6;%7}_f89e0j3IXOip@76dQ|5+&n$A^#kp0^(- z!@W-X-%2sue(V0$RZi{6(H)C=o0_(!V@ag_vE=E zU`}$Lm#;rkTB+JQy{88wK#^h3G%*aTpC$B}T$^f7GLvS#$BF;x3ak!D+~kpxrtK?B zvak)dP(fySebf1D1ogD-Zfl@77`*ahZ-CjlB80?D%4`Uc0-4{^7L1=MmoGxk3jdgP zM7)=uyB!bU=SX$V^gWI9p-&}fp$AEwD>n)=vb;1$zjJ64d~~E;v^-Vs<*Y00#dLCX zlDzFM-F6!2o%YVJrp|b5==R8YY}>q9U> z>6D3P^X&pycW~f$)>j1!CiSCb5d`{@EGu7!zaS-ozdZn))K!>Bh)}$7NU}J5##UJ%FX~=phc!CLD}Ic$LQIknxeB-iOvurZX1h{j zkfC3ui&2)l`w_qmzF>O=u>kv6d%J9xF?|*M4by-tpDUeMYP%m;0A_`3ldutw?^vcl zuR5uZK=-d;C}V$7=TM!|0O^Ng-7FuR!NyZ2#=J%Kv_D| z0Q`lZZGh`rTz?87(BO^%q`}gmA-@U$meMXgn zHYW3OQMy-;t2tP^v*0I#udODb{&Di}fZ1h5N(}8+xRT}m!Ega@`bXr%ug24nzae$Z~XJ&K!o#IXECeN6D zsug#KsKKy;wa*_hJmn;cv?aUi5GKG4Xp<=Yi4gsFI%e|Kv-{$mC*<7NC3!Q=7erae z&)8k$+)0OivJ7v#pbb?$m_`^65tI5B7Rm6G+Le)vOEFe8p+81t5$XtKDpB!arMTRE&ae)B^Q&3}(?xIHWm^gD;7ILMggCNJIp{O#-&y4$eS zS@#w(ep9vBBJGeVgB!%wyzdN*I;IBZzwoB`oCvl7bLpil6BUACpq$Qf7$_L6mW|!$EEV{j^1czg zIW-}K(L(1745M#OCV;>E5B@RJE38Me6lKi<6`4q(*zoGM6IPY-cqRc|;hK|GYKR5J zQ*m{x5wh905QNBW-2LT|YW&D$h3XBf!d{Rx{M=+W5NXvF0eniDnSwRfnU zy}ixQHL+123@YolJ3${u?lz!{d!OJ@Q8(M1&Rw`5L>y9SpW1yL5YU?uSe|i{zb&;8 zA_7-eFE|vuomlnO2sCL@c?~LX%>RnyBR;%4TVS{w8;@=C)%%eJ_JbgFs7oj20vflt3`+)+}5nzj=mj%FbGq_B7C&i|L$uX^#o;lm;d;4A4Ts zcbGRl%j*IDB$p8q%~gJd3~0 z+`xHxZ?_L?0TVEssTD@D1|tDLPU7$M)WpPxx{SKY!w>M<(lM3NRrE7wnH<;e$y_fX z%YI@vX6ZnA`CEccJmv@V;&f3;llIlOeqr$8PcZG;!`Xxz;Dt)_atFE~1AY5PCk4}~ zBqU?&oDuqdR@?|7hdG%vlP(p;$yT>PTDS+2;D3>h8%TjeXAD? zH+hj#tucPOd%7mkMAO0E&@|s0wYiXl>YI6P@#DTJ61g^a2~Wv<6z^aHOD&Do(z07! zffT1%3fI{Wib7T%?5m3cJQrCPvQ@>VHN+%(*T*?cs=y&~n)3~7-y+b5W-G+y4U<~0 zG??!)59iWVD0LShyqBh}p^@6?+uju-=z_AffVtJsbz=_nx{wyto2*u`EEl4QcS6Pt z^mhv>jYxowepE88mkj1)`i!8P$G`rSyc?;@4D>292mSW{&;tZUTxNG(j-Y^{G|S2Z zS~@xA;8h@p%5kd{v}OI1V^b2uMHB3eHDxA%yfrwkc-2H2`6^~ z@%M+Om;l7p;CV1-XHx;6z9$(@WrgR_Z`13_^-)t|?< z>kUcnrxP~6r1Pzttz`(b{d{;cFhpHvU*2x*Him<31asnh84ve$-CgJjpRAnR&hZ&o zTY)KG?YG||3JV{{<=m1yE91JG(YQTDPKfuE^87n)knCT7&M5BZ%v}9FN|BG> zBL!=WQn*_J=_G{rR-ZNr)#$35J?7JeKmk+^dj;h)aig--^X)gO=PT>3p4ScM*pFwK zQb4~(Ip11r)nIcN_}$PLPC|~ogrjq2Ki}#Vl{6dw?tcDC(*=?_;~xBw7#;TJP6UA8 z{}ir!XMLmsS5n$+yC&k-p8*~Tmus(J;N*RxF(Q?U5p9gO>J1D~LrP#Lu~LCw$z}(5 z2g~mfF{NS%v@#b+-sCYu*55;bK^_Dn)6+hDz`%>K5zZR0`{k#RdCpj%U;qF_e~)PE zv6rro!Q%7{CuP^%Dt?(gh!{vLc$fbRb0{HAc+a2z!oB&Xp52Zl3)$<@4=OuOM2glh zHd}6%U7nh)8G`1*K=-<+nR+57D|H_0aS0f=JSFsyAa z7?8pYqo9jyJ_8gU`NYrf$-+yMg!gA{2_B2veb;-VY*byTu&KglepPY?o~>U6Xi~l) z`Gt!$%~(CA;RUWTjbSDN`fhdD_=(Y@Mpo++X~)vsEfLpMO=sN)HZ&KLq`mf#npbpR zv=(MOK-)$jld^~se|Mm@W?yfw)5VMJ-zT!-vKz9XXbbZoa|Ic6C^U?eS15v$$D!n{ zp1e|MILHE0!DT_pVFqiIIBU3@vcScf8Q#n_Cce!^cUZTl0DSURQqJ8tg#mg?pUwi9 zywAaAOcNq?=(qyL#E@)!Od3L=!+X#EdLQx@nss)jkBBRMr;!F!2E$Zg0UWwIcGYGo zCmH7cw<&mYyVTCwu38~eV&NYdnGg*}JoJoqTm>dqVzlJq^Jjc?fXM6-siXjU3pmjl3IfO2p^ci2mFP z(Hk}Ui>>X7!2T)%Q~jLj@sU}3o%8bYLV)lBnj-uH)jKz`gqd}3<7RWim2$>I)APmc zi|TJC@S|T9WEK9^E^v|ds;?a=GHzL2Qs`{eBNU~h`}*%uFpSQ~O%+DXeFCA%n*TY< zfDm0A$=TEkJp_)jg)9XBZuO1920r~Od#O?bb%RC8nsQ)z=VJ}7 z^}W6E*VVH#BR{F+II=1=Ubt(%Hs}u) z56OUK2Q~0KEi5ep!=ZNxoUE@kI#GG+Qf-fTFp7Uh%6w+yb^Eow_hFZBvzzvQOAjb9 zgm+e=q~z21W5-xU&dBwChF4!|eiF$XMn61cM!3Ry;6!M8gd6MtATw|}jmW~Db>x+m zxmHS7zNRYv`Sa&Py1q!=P`FS@=b)3+%*WR&m;POxHKQYQKi>YiE?VQ5e@`Q(ylU7u z7wdfzd$B%KJ{qt-$=Ml@tEyacu~^G1JZ3L)HXQPrZWm>y<$koqt2SML>!5qpp&JdZ zlybHC@I^(6xWhPNl93Jz3#*J;*nLw*fA@DmP@IDHSQGUm!3hlc5@LC8T(;*`5(EBi z>WJW-Y+88l@|c(k$IsjM?%;vNh6Q!x%$u7!vJlBu`kL5>akkuCuEuiUXmD&LbQbqLtvaf2g3PZA3bO^VOgfwny zaLNW0BqIA3NshhHX4l@>C)QcesHSl)0r6Fg(+}0jCMWh=ND2#)XLHD8vA!4LoZCw2>5gIWP{qBEUzTSWcJ}nU%S;O zISpR3lb2ZVOBQXE!=GAi`Y)GIC8-4+efP&*3AP~ksA$a|dBE4N<`M0TyKYoF=IDf z_0aVpxhE9G4GqbE?-R{2(3PaLJy<9uXH)89^(^t8f5H|ETbzu2fMF^#12bv-gDhCq zR2c$OucQam)6$LPL80m}nKBwU4vSNiZ$d(EQ;E>!uM60%-|*}KI}8dJFhrUYJ&d4H zTr%en`N?cFLRc%tEUvEcgxNfc<;N#*1>_^Yw_y)zs+rgTOhml3W zi|b1NxDa5F#+a?!?&FxIU3PDfS@9rhFJsdj;PozEewIix%gbZ83X7t;-IS25AwS|x zBErQz1Po&rND>19B56UVIi8V<>l=&lmpGd*NDH>Gp2dzCcw~Nuj;xN6w1T1x}fsC*=|NPLp*|z7Cr5%EIfLwsZ>i9cj(W8859x+QMb$j^gbe z=2kS7ox(xCI++D3+GaKz7Iwf@bGt3q1@K=Pw_fGEroDsT|5Wr+UtScBoSzTYnM%$9 zl#AcRtNY;|h?DZ;WFjqg;galyL_ANN(qIC6*i){hhq2z4q6M96Whk&X?WBvqHekmx znyZzo9>8;vr#Ou9Q96BcQz=DxcIyMEu)8!X$mEpuPA(W2dh~ZWs^FWv|0q>lTpUf` z=yosY294b!GZNL2W(gZarSu!9?`_;=_JT1_??gx(O3ZlRM5mLjjuF95ib2k*MKVJW z8BTOPF&hrOuK~7wO*#|Cz}FXYZd-MOCI?M|oR&!q)<~%OihG~w#*-}yM?hJohTb{d zOyMZ+S)Vz@*c-xqAoBI*gapBI=10T|^kM`*f+ZXd3JjI)g?`uDBIl{Mw+jnWWWbg$ ztgJTdaVD?VXnL4=s$Lc>j5)sduCokpHG=`vm0upx&7L+cofnX7{74;p{-9aX`irW7 zY)0h#Hs;tnLc+p&Cv!=dhk&c?e)%&TIXA~~dG`10>;&&+T)GZk zypRp|#PDAv_gNz9r@$_L8vT;ae~rJ)b(|PMDfz$4&)fw%-T}~{JIT13x`b&5#uEk@g>ZB$!*LQlf*i*H3XF=DFX3oNbWZ;Az$SZsptoZj4W+F3Q<8 z^9p~>CiUz;j)X!J7J%_DE|E}{x26)CBP29xHs?3=Tkwx_f3@{RV;_HbI619KRC{fB zs$%4AVP|o}Y2w{xqMtOa5k;oQhvamT~owgN68nF(Tb=bG8_5)Y%EeBY8%whUjt^g&>{}#mSws>i74*wXV z^oCqk`wlD1BEu=1KSR6aFK}W0FX&+<^K$ zi0S+-KmGTBdX+Ck$Lc$mnJRZI4sxOgI5*gyGU4h~;m{ID)ZwA3nZ&RoE zqS$4zHeXxi6%{$`o|-clR`8Ne4(^mGODd>eXD(~Vs9N-mphV?1az%B{M-6v#cM64s zQXl`xD_NenNnpRN8Kf35_O;pc7?;dF_T`Y<6(RsoXdnA&LHgjy5+?ntX`@2>Q&y^Y)} z>GA$FK;!u~oIN}z2UcdKGW&T47I#blZ1sL042MoDoq)WbKcJ$Hr+Ej9Q~lSO3rZPd zk?r|fAPs!ksBmmC;y*Zf3moqx8PZA@!p}}6AuKgtTTQ{{a_XRnn4h$_6ng0&YAaVB zNtQ`cq^*o$s>H%RCCwy>vu)aJU0wvUOb(X5O)3KNZr5eG!zaCVcrbZB4j=ch(xR(( z=URSE8MXCIo>-|?Yka8uI(R)zgS%7wm4UW$ z*`Mfts=6>3rSRw+Y+D65G+l&UXAhXv4njsoVnyUEDV}mq`nEz<+%t5L2JdgUXZJ$4J0Yma%6naw)plpgu6;#~!h}OZ+y8vvq2OEodGuAo@)dho`Zf>axGS^q82F za&>K<=&1rKYMbx29p0!An>^<*pm$O5^#_-Thdv*`M+3qJn`di#{4`~{m&lcq zMLn#ul{vQa^zQjw9Qpo8s%uQ3;@#OY+a=w%HgdW7a@OIxxY?J&5{9$2B+x2rDcUQ# zZ_*TRphrMgoeCPGP80Q=`B46>!%U-WE{cpd<=jdd3f;W1)ovLp)kn5z9A|s{!?^;8BidAxyL-P!KP z)R%uv4R9-N{QbH2F5*4==Y0)KwxXLkQ{=x#>@^0U#45`K1zxpp34OhXPmsr>llqc% zF}e66U#%+#;^v-2eLh&i|MsGOzTdrrl0}8stfZR!lVua0=$Vlv1TqA!t)yM9xmLip zz-5Nl^U}=d%HOd;0#Y-$;bb5SVp4o3*VVg&mz?n~XGoag0SE-hJ6g;`YEysR*+D2# zbH5p@Z*hg|>KET>H91{B_h?2HSenYHHGKX8GZnn3~7(kJmFSiEZ` zVjOZZ%b9I^u@zQdxTZ82{#7sP{`icW6?(+acZLF=!p0FIg`t(;dR$Urc)O}JS=lUx zKDCNwQl5*28#;CsCnItno_ANg7jgj^0)rX$6!|$^bY2TDmKpYlLOja4uIp361b7)w=8blq@9MY13_3y^=tB^Huqh3FyYi;mcHUaAX1yAJ7K_4cUO8p$Vpj0m@!xI46j_ zJ&<^NXbJZ1kd^Lu08QYeqA|8O1O4vD+Wlbt^9PV=dBHpIt-tjy5v zN8>~)`Q>hiybU9llI@AVrAnJS!GH~iD!sTGpMaIZB;yq+fx@gRs*)SGet+H)b4SEi zB?u9Rlm2-fB)y=zxQNnkqw$`9*Xw3+IOO3|)*|0AD$Z=XtF5xxs<~oJ!0^?J*dDW= zk9OW(=w)F)J9dLeeOY{icneAI8u~pxg-oO`FRWeLEEu8jTk3U)XRmMFc#1}<+aV zum9`eFeH&m66XEafBpxv6CyN$5ba|Zd41a3ngq_5Hoe|-yGPY}7~~!SD>Z3QDQOUm zVeN=GcxV0=oq=>4H}(FsY`gKfbl%Eu<{`#VTcX_&Y;v^%Y+o*yQT#MSS;-ZW;y)51 zp&TIYHHAM$PDy0K(xmg3&`K#4o zxj-&QOgf7YMLC;b9i z;P9>BhP}OijcIz86?{aNzYU=QwQZ-J2sBSw4{#8YYH zApKp}QOW`B(p@>H?2>p#O*j#f=z!Klf!T4yCH*In)c#{6MFcleN8v!(?lFH`Jg0-D(e&tQRY!_)7 z8=rSXWU#Hdt82ZuyyEYFKFE#}J>*C<=DboL6Q-;;o~W{qZS(i+TuIutF7)Tv*+3n- zGbYSWzp8s|nqBvg<<5rO`Tkhj19>g)q?ft-S~v~zW@8<&zvbBDu|Bn)Y0>%<5aFsJ zS*15s?D}=Cvp~P=Bx*`ZGQl6F+(2$}#>v{U>G1Fnn3Uvk*5fCp6WuKcU)fH$X!()4 zfie`zRPvWAKXL8sv9tV-mxOcYC3O@7Dtl)-YPt?o3?GgEJFHrgOP5TBd)EJX^B{lu zar|Ki;EXC0P;>enk}Cs*)TUJQhvm?1Iy%ozDiozga}xXm!-Nvpt@!Q9Mja+0PEj%~ z^Py$N_U*lvcc66o6m)9@-SNBoC~U!Kzt=gOwM!MIsIKF$+;i#~VeuPz zLK7QivRE+XD0U%=kd~T`ThqO!_Kdocs|fgBYWiH@H--?F``^~76v|tj9tqn!ZBM*X z8p^g?>JSsA%&{cO@rbr1%&}c^trpmw%}z84xfe`_`B38NhrN-C%j@gK1*4h_As0&? z&*n5sc-UU;>Q{EDodLmR`-|rIZjGkHp+;R_PlxP{)3BW#Z|((oHm_~pyvdV+#mVY9 z?#a}Y-KQFHtz3FevSdEHT15P_r>*Vx;N8ypE?)X}2iT{+uJfz&H>sm_swT8HaMqd{FP(u{EGMKVVa}N~D{1U{-65`^HA7$7%O!d&dzC?5 z?b8}12UCnE$H@0hEpCA0gpoeq5#@$Gz~cJ=xM~+)e0l3ku(5T@gJJs&o1YS`0^%REO2wqbQsbx7#X6SafCHTe)c^X)>^%b8qr0y56v0>S)Hq zV~bZTcWf#oFWL9AwETevjm?1JezEK?1UEjSrX$82el2+tN8?v9_3W%$JJrQ{{zTj1 z>ZEBT5Hcg!zU$rcgC(o@`dp0c?ZrVs6(l83d!~CbJerr8Omx46bJhQe?<=}=D%|eG z%*>AOH8l^FB_tY#He3_w=y-WYmoO)ba_#Oc zxf#LJ8fw0F;T@WF>Fy1$i*Dq@wc`HayU=^9C|0_h_|?Us8B!)J*iI6zcTEr}w z{NK@57Zr(UWr`?lxQ_P?-9jd77;P?fMyV<{`%WIFetS+vre$G6{wztgOcd~2F);9j zbXKNhEP-6l8FHK}`4N=*Ui(L8IZH@mW_Qfb&3=z~k)4}lK^HSe^A|7n=g+XJ$cP6!cGIMDw8~R=m1_V&I*ClkaUT>Wu?8EF=dmW=&8SeV)UO5)Ap)w_5 zo;oW7i{CkX5A1&xoP3P^WOuH3G9g(Iz1;^oE~6wP=4fu%>GTq-E4tZl>Ky6UJ*%LQ z5!s(mR#o2fFyGtO-AmkX?3B!$Rr5Pd(eQaOkHlk+ORJ5@+-yF9l zr;y%C*{;;fyJskLzGRbRc+thoH`;i7`i77wtR8#Kpkg`7Yu9iACQnZ3j|x)X>djxF z&AP67-1;-IuUYSa0|5WOBrM`3u+gR71IXMQ3xC;&lezsrl`G7;M!y$3`alc#Kj4_S z*d&^Q`zn&5Rp#9EgWgcqcE=sZDk#|2Zq+e6mXnDI9jv-lh|8R58Cypbjhdv&5nD>l z_tDX8n`uG_PyN;ZCHZ>%SGcQURgrV74P_fMP?3Ok>z}vD>upR#H6XgQ@^Yw*9}CX&em_`WS$q{Q>CGt9EOL(^`S)N z&QHpA`SabXD?akW*(AB=xqih=91gMrda-R%Y{W?NQ#dw1e;P#n{0@14cWwX`nxM-5Nc_@!@Vof`J+o8X?{_3*CL8$)4uTS51M^Xrrtk(=v9}x z<$|a~tWB4k!DF%_ABa8T@1&Y!9<%hpW3jh-PYQDmgP{ubhijPb78{&Fs!g34Y9SX+`F`!CCriJdJkJ)gX2U8BXGVpI zG0_p&WIAx*95bWFsU|e5*6B|bU72&3Ct}fs z9uRLxZ#Iqhx}u@h^oxQEf94wjx2q+2YpEd}o~~WncPtc2IW~ukTVzU2W{ptjXW{AT z=MmIzhaDXq?66@G)I*tTz^(9wY!aT^!+Z%pY-YOCfbpjUY6fB4;yis2HGR*`pb;g>%j z$PBk%$M7^YJ-{$}`X9RsdIcaddNMXtB=4Cvy%59{1UIRo7&bbTj1t^5T+Whcp4OK| zRGO?zn%hZq;~&E`oNJc*)egI)%w0OomyuPJa&f0_20`PmxAreArz5Ru4l4Qit>8mt z+AddZ8ou@DkG;u_H7h=T^JCX-jWpAr>s(QNbBG_NwC?(O%HpVl5oR&EO@>P0WbY%|A_M!%_`z-174V z7*trrgCkk9Maef&Q|ye4jNKa_;`GVS(H!})EGaRW;G?MntGEKJq3WW#O2!$F?*-o* zOEjQ8CIvc+eQfC#jaltfaK}!_?eD|&?>P|rvjcroMx(XIBHB-!-LDr7ZWB7Sc2eC} zrETq&0By1y^qU`-gNUAuDiAoq4e`zzYJqcIdh+6uCOiZpaqqX*tp9ABSjmTUyiY|B z+)=G9iRlSP0A&mReT62nXl@4JE?f*b=5Dn6`H|o!$He8HdHg@Jo;s?@M{N`pl?J6t z6cGUdX+|kS5CkL!GE%xbq_F@==>|bU1|v7Rl#-4yVw5mC9i#Jm@!sFP-t#?w0O#x+ z_Po#gbmk>2TRpsD%yH7R|0S7RRxv(7diV}IGy4$53eW7k8>|ql*JNw68o$+B22u;) z;Z_`>_jXBJK+pXxJW)z1@*V+SES0SzSnK9vJ_D<*xpIfIL5=m>t6S$a7$=EB zACT~`wdC?$%S=@%&04P}100wOIggb|gOMn}A-cAL3pGB{PEW_A2*T+$ELAgZ!Ff6H z;)wFC-5YW_0M8Q;gz)QxEUl|G4a_F__*`z*l*EjI|j zb>YLq%HVtBFp-7O{c8J26rZR{C&|wKB<({z{n?(!Y3PkNjr#9|Bxv9bBQ7RJ2zrR5 z<*Ls)g5xfS<9el`+OUNEd5cVv;o%R~s)<`FX}C^nSBILqC*T)_-p4Dkf_qYX6OZe* z2||iAb0>IHedNI;O#vn8y8O2!M27%@CZN#xgXsY7@K%UTWQZu9i&s3lkAO`KMdFAm z5lj;$3VdbuutwtNSncMGZa%&;QI(mSnv)CG4omg^8JiQY(2u2x-UgK&P4T^o`shmM zu1g)t%eg0i*Hgo#i?qqx1DB9Fee8W07WIgRQsJ{Eb4R6QZFNKC98!Gi{75NRYi#Dq zXNXXMcKe59#?->S5(I4PtVOL|4eAB49Qs!850ALU@0&hd*TAhW+4SfE`sIMHx*Zeq*P0@{L zz0LE3Nb}626_S<(RT57`&DQbi9O;aok$m$Omqb$@Qq>=4n>psW(a)NBdy|QPy)dg@ z=#ipzCj1HiJ*Dt!x)bw~k?clhMR&iY)-`f>1qqUMf!9Yz92owyZ{Yh17u`D1x!Aku zY$TENCX*gHJ7vAmWS0r(1WbP83t$v*$HIp`la=-|Aw}=5u2EyX!|RYgr>M*d(WSW% zp%Z4(Jv?sKu8wYzD00JR?q$yVW)@S+$+(r_#$)gbxx9q(PisMdKt#IxU?&jsf)lhk~85EqUc~y?5tZ`uh04DW%juWn5d3U4-E6WXVz(nfI@eJ z6gi^Ou9OIXC?yKAomgMElO=d{B#oI3HZYD+l#pJ*%;X| zLWEKgNK}zca;#-*UtLvv%tGez-cTEBqirzFjYbjZTF+F|c?pg**mM1+B~uH1;>+#L ze)QEsoyI*g=Z=+>N+$~f3VP0c2U`-2#}USc%SBh+;qc6V7I3)um>y zT6zOV0p;FyNs)_!C9be5a}`9mT?6y+YVx_y}a0Px^4Cu-{z=aZDFJj2sQ z$yDqPtg)cZSN`?t@LPTSTy#Pb<2chpdxbp~8aN7W6FaKEQ4Uh8R%u$SY}4_@i^eO_ zH~0}|Y%dx$3S#GrX)A*Tl{b5CU2vu9h zcY*<6H1aejz814u{WuNitL#ZJ@ysyBU>fuiXmYBVlp>g$&8JB_ph`bzY?43?3piuc z>}Bc%1L-e$owMwYeM9tHXKk8OL7+`5zWg?3-?@M#=9#_L5hn}X-)D}%h*Lr@NJW9X zSE}JW0z2y~O>ivC3W~8X7_8e+`O~NbSUv)RvEU1`1rL5z;I!toEwa&-7knsR@a>h_ zn&@tFo`_NB@#6foYc2+b9Y!M?g0U^}w$-)|CsNGaC*wcAVo6NNccfPv@WHf=i7X;uJDR7ZE2x>a4gTw?U`L@rY!)EoA*T|$$G zAZRM#)CTY(uev5?mVu(C<}`+e9M4V*W2wuZZLBW{bqwm9SFVq>D^z~q?m!gEG=k|P z8496luGrzbj2Y5Q#g>GnI?5^#q!qK>Ve}TietvRkbeqsWgC^&n-;qNqAaLEa3VvJjh(CzgWz`PsuWiRi_OhZJD%)@ zRSa!+>{qIee)Bz4Q;u}>KiMrcbE?6HO}qQJmmZB|2DId=2r@-oGp>lO6GdO+vd-0a z#4rFEzkfq1C7^Exh;rlZuc@*BYM+OqfxBSZdM*FpLqH&s4IEE8`kb6n zs|a=hWN5??TOl5qTYIgag;MHh==lnWIjD>EC-Op1%ytP+TBAtZ(xolNOB1CXhMt&< z@o?_7Jv0|%TH6XR8(&T)7TIx0!8c>iHHWkg7taQWUfPV18eNz0ckFO%%?Sk^6>1>8qD@hh=$exy?Dqwfu30`(65H@R{PC z6d^y&--2$dkw!D|^D9=?*LN!{_4mCjOEG;jrvtZFmoAm6*>2z%;O;^0@xV9!1Fa$# zRG_7&r^m{h1pyj6J1o^c()Dc+P5&CyEGFcq0?j-m^k?9GJksr0x$e|Lx=?%8f*SLRp!=**s5zF;HdDs(SW@I&AgeV8w+0rIs21WRF!7z zA$T{$ZOGOe;d$yym&UO3o4mqLl%Td)mH%OnuK;ztO;f&WuXiB|9q3Ihgz9!MYQ zii~^;xl#+VY{aecHBqiw_E&GZKH>arf4m@CCwD4Iu=mXmziZsv)B#)*c(0=T-+fpn zmoMao(HvSlCd|}FolG6ZInKPKk0hN!D;3xe^^8AOA30du5#p_~c(V0(8+XdA(9L zo*Az0?uVX4C*Gao0gcWUeSsrup&?fg($0u+K>ydiFb&|RF>v04$Quk_T0uyc)vq?oGc!H*LJ@gVloAfpe~B!2Ds!O z;JKw^=buGnk$ia>~>HFY@j`q`I4HB zrJCl{i$dtghm~M?uI^8McXjoZphrinGG?I~aBX)Es-`thBj@DC&B&Khf&tSQ`JfiB zNtj2i-6)BprO$wc48dUjshdi~hy)5oGhh=n1$0FG19cZl|H*?~9-VdIpPAAWWxRGx zs}e@MF&xdF$Wg3>rT#hlb^}$*8lp=I}bIk=atH~e}nuS*Bdh=fgbX-AYW?f;&-FP*IlSHb+$-Z7#IPCOG5J z*5O+QA=?R%tdROCq+&1y93C@vWcY^^1kL_{ISl-`vomkb9%BwktzVR|2Z)jGuqw_2 zjp;vgeXct6q)JnTe%TLNdzw4tpMDGDR`z}8*w?@k#>==j)P8Hzvv%^u;#v!gZ@TL7 z9bWD+ElJ4H0QqB;CeeFq9;7<@8tv6F4HwGI^2drFDDuOcIKhvSMqi^)cX@&M|7y0o z7ZPOrWQC2}L;eI%u#C~4JXlJ+el7_oru(e>xjsXaFB+EUR0ZcmKco@f^fxcDO1FWe zdLKr3bkHZny`}>cypR!n&_YMP^|kh`({9hFBWap$hsXH$Xpi#+OZO~qvOxYJsHdpt zi<=q~_OAb>0vr290)ZH6)_}%U$MpA8wZ`a|}FLfZ}l zL}l-y7(4sW$jXe3O|su+nQ-i54@M>3;>(*VvTM|6`1Fs@AH(iIQ!NyBXRWo7IqPYD zn>q7r+-$GSyD)3Cx$_ZIG&$>&`&C;X?QEuj)SL*Ec@&zjUH|)y^#x~0s@CJ6_Mc-4 z;=Ef%`2RtKK=#8QOaw@7RGR4I88~-~fhtbIs%pjs7f;@c_Waa-U!&_Kr>2;f)s_ut zcV(oAf-d^Lk+%ULX(r#}XHK!;T&)TSyzp=*CNI?BSPqwVqq`)NlRaG&N*dm6%cNEpUtx2 zHwu~#F}8iioBMY1rY2S~Z5gqAXKm#xFTj`EbS5RIM)(ETffu5RLvImdN~_MJSyx_4 zc2FiJr6n1Zj>1N?eSCvvP}FvQj{_eJxa}B6nmpt-a0q!G#+7O5WB_+rM2Z z5o*_?@m~&8Ce2?rYErjAARc2#M3Jg&1>49uMEA~1!6X>S#=IHPS-|oSS`*>X| zt{Hb(qRXam8)_Y+w$P3{fJPb3CJ#7tD?OU!3>c>jU7GK$F^debse&OI`x& zZ=7x1W}hpE(FZ*Ti2e^eBLURYFD9i<^hgH;yDYA3Z&1O&-y77I(fke{MQhQ@8sSNg zy5`!?i{40=Q=>_>$vF3u`yBH~*Yns+uch_N9Jiz4x{{)#6=cmux}mgU)@=7><}Bce zQdi)4b&e|2iXPrL@SV<*CX`)HAKQa!e%A1DbzQpjw`r#Y0~nozX>`HmjVvI_tU7Lr z^kl*5UaJfZL}2CZ!w$#fc){ho{H3ncc5+^b-4I4dk9+?G5>^u%+pV$P4#dD*iw|-a zr;k<^*&qsSVF6+cp&eHYkcI9VyOBD&NMdjn>bSDwDpx}ml5FfRFF1tl;_;1!hK961 zDW*W+%jJxAehoW^gq-!$%^HQ6ckj%)A*FuPs9{c+p3lwTTRhK%6CK7&R9q1d`=Pc` z5llL#wu+JGnn}y73rJ#PNQTZ z_0Qq7KQADJ=g-&2-F#I%cAK*Q17yttlbl;T?xP=8-6tc>3&NFlzl;T_-?D3PC)t-K zOtt>EvF&FKjNGMKgC%Yip;Tc0_X=!>iGd1|CAK>JX+pybl) z;N;>UO?S-Bk}f)K2(;b@1or<1FcqMQof5P!3Xx_OiM3R*k6w8)J~GVI-N2_P7ujqj zi?0qFZyoEKyAJ)*YoK_e1^gKoyd3;I!Q+Z_fZY}~B)ZS1v8QJ~nZ5lAP@*X3me}YS8klr_ z3U1d$TP$&?FwNW7RiP1TtiG@@bm>~IC~}0j$VPsvHUn97Cm}|UU*Zu5UPa9Zm4>&q z7hXtU-6w>U{wfFhI1CX*5PRQZlf_-uF0wXP0uqDWJyNMBN(z{Xi8VXYe^Kn3ih*xg zMXB`0CXhaejPE^77}vz&(iW zp_+AiiTq>my{|9tak9<5+%AdM>!Qe&M}$0}q|BYEE$iR?6ms(R$b=z5^AHTq(|5-k zGcI52SweNxOE-+cL_neiJ?0O(Z@BnbukR&#`!4>-$_-Ny<~tXEb@x-r{L~ifaFgSX z!`v#j9=Y^kK5=B~rGTt_T@t>HP;3n8;|H6k+4oRK>&vADJ)p!-Zh0c8C{MLLGe)wBSQI!JBilcpe1GcL zfPRt>%91qS8RPQ!_|0)I?N;>DLX3)_0;#vu@}ep4)cXdJgk=Ac{EAAK1PCt|XsP!g zj)Z}&#ki}knAH;9!^y5u=h-rNgWHiwLohQ$aqio@uj!UBzG^irE>u9SK{RlH|I|M- zzrnhz33E2{B>5V^-~RxeGhq9JDA(F0`!R9;3x1|P2EkuVU&@Mc17xm0x)_5`!qsD! zu1WN3897bvc;9yu7Dn8&u+UQ@Bs5M1B$hnv8@)C%@TW)JMZ@5~AsR>;P~6$F>Onor z5V?0gDO|Z_*Mst6hc$CidA>HCZt@(*OeTq5he;3JD5}v!c^*3N0{ndghY;Q zz&lq8F(>clHRL`mT&Zb1c;UAWF3fjQoHfzv4GSn^LYByT=Nax!KRG>qc7AUuqR^pG z-1~{3k&mhUvj>|pus7q&rg~@A5DiXGM@bFDUhQGl8?wOfpuHpF@AvyWs*uPHr0~Tu zs(%5+Re;t{e{Hl&Rz)xIyEmp8nEE)FX2eM5au`jgBqeU#a&DMP^7N>i6t}gT8Bg+@ zJ*t^~t^oMcN#)bqMrAExv3if2#v-pvH%Twx4tY6SKZQ(+-K{e(-IsnfgWRG%pXap&l$) z8U=LGur}}nf4DzwvVr4!`4{uVk(8Bh|6ot>hgJFn{YTLBRQ%iL4efhr?UMG3*3l=f zBV!gJijkVFcYuSqs{v|lpe6x6?~aA>2S@ku5zK{D;M^jhKz5r(Cq$#hW8TUR7Z>*P zVPQ&ExU93s(3jMknI%@wS=>tLJhl&I)~Tc8wUIHl(~hnw(Uv`vVL*ND&2VayMh==cSpwZ*oNNz=lCX&piQ<~lC zV3=@a#DllN%4eX_vBs$+Q4B>s6~|IhS$u%#;-{F3b7I2$uTGDF(IK}RH-xV$#kRb@u$j|Cx(nD)rPZvodg z0LOTu;j4feqZj7$NfkiL3IRT~u=y~>1kF~z*QOc%b)+BK>t8=r7&Vj5jZ;!nuED?( z;-qSaXCujdsMD-hBOj?5Z(f~xvGkztd79YRv173V@pDbBqJpC(zQ{1eEQtM?eP(Z^ zXz$N?w~mf54f|~O)k9=PGi4kNsjw&N0fy=+YkD>YkpJxr&^$?_^}?W}vk ztf6b)juw2OP>dvHMKesh8t~?s=~E`^>*7WV-=Xst?y{@ANbQ0U2^unXbS^N(s~uIbG(B>8mbUz55@Ntll;x5K!KJss-iWM_A_){QVA<>V8a_OF%T zH6y>o&xW2bERL&c^r-*bQr0;w^dC-@-2uFh1l9VpmsKHdVT$mOxZNd&?rY6>oz2FI zOZ>{<$=4^QY+co9eMgFxV%*%^&fy`h#Sq?r?@D?uB2_-rK7~uU)vE2{5l_^v`!;5a zo0pEZ!(??7^!8V5y3aQOQSpf>#{=orSzf?6cdSH)k)ux$^wq{~#UUXc)lWxy6O(j* zA#E1dtZji+otG&*Ix6euZ)gL@uVIycTLpPQL;WPFx%5(85J5I?Ew^n=5Uz4onHL1$ zrA)-M)@s9(8Sbsb$nu%l+$=D%hE<#_RDKIad0n?nGz%f4EFt=XSm zme4qeCY>|8{W@vj;Fkv2Dmp_kSgiXNvJVDy!Hd}_-0o%A5vo74ali@hy7zVA3KBHC zkG>Y8bQbYz^6wj@wb2f~Lfz$GwAO~X9V*ktM28vFpSH9_%rN+j_^03qDq02)pAfXzEv%dap#uR*9i~sh@-*VB7xIk^20Z3|wqN!1wKgxc zH99&1dM($dPdx7aItw_-Wb!>ygV6e|rbfJ#l~w+1(b&3P{~6w_(o3I8Uo-;ssn4*G zZZ_ofaCk^d*}Z6uGJE`2+zRBU`Po0@Tzc6{oA# za(+BMB;}oMfto=am-IzA@W6f!p)pHIok20@3m%6Bob+2$mZe=2-a5@B$NUzaOZHdz z9BO)w4R|%K;d6rI-a5pc+-O7va~|4)GI8JafL-TTwbP1%#ZeC zoTQs-ke2yYH)<2K!1wxI@66v5Kejs{&MPXecm}pw1S$05$v0^R6Kvh`jog*fLta#f zYYgKNNFf=Y{$p|>*1~ny)8Ai9>IKeZ909C5${%aEFp*)+*or#} z;20r3tHm`3*#5^OU4nP`{T(+i5T?k;5s9&K1wylb1KTQK&A~}3q=AR*F)YP0a!ucs zlX?=9W6H;Fb(XMoleNmzRr&dp}MxbIV3&*gr>CzuFwBHam?gyr;s zz@Hqo>B3oPe&a#o8^K>&qcBUO=K5czs$98^r=5)ble?&<{4m9)Q5{Y*DA_tARm-dm zaZ(1c!{qod-DVS>pilj@y^siO_SSsx`bF~X-w;R{Wg%@$chi9Iic~RJ!zA&~qxG5V`3|`1_!~M~=OG{RC@E6W`-=aGY6u z*0X>1${SW@?i{e`ZgPy`W{g@N7uwn@Q3+A&Ua2@l%8=;j>ohOb{Ho4}4_HfHj4w$S_INW|CUwu-}=_o^Qf$i!dtBtKqCsTeh$`2y;6%U}vze zUboOFwel=hndyhD#i7B}nzQDKQnw8z&O0bw-s4Z&o&xqll zN01!`JTBRn@9hb~w@Azm2`wsLuiBvo#sT-_^|i_%KI`+G{CqVnC7%3|2wk`5nqrw# zBP&y|-H)-|MP{v_BO8yKs|#Oy&%PTnhF#7m?%Tr>0oj){+^67Y{jZG0%N*}*`$$Qt z`R)Zdz<0-{`&=cMI2Tp8C87FB`HM8Pd3@;*p_Z4h{u`dQw23`+Z=RH$a|ZZ3w9FV` zr0ZBo-#^daYp6VOTgjuAPYn3IT-~Nq!Z5(PtWOqf9{ahZ4&*Sfe}$h7D4LKN=>(gb zB1g(Q|IR)O0hLx+#nZw|rHe-Aw2-)=HL%`jp26<)=OR4GkBP~-s#?2A$8;^m_)TODD)S%rXUE$AhOE(IoKnWv*)w*V{s z#LM0cpne}|%najeX{54bmMl1H8HkS4+8xd_5*%b~P5af4F)d$th?U%XYH=;D_L0;` zs2U5%*>|m5`8_(8deXfAY+ZGf)RvY{U5T{j(j!4(F8q1p-r((X&n<<>91hOSjeg;3bw0cpP zRXdB$QB%MczyEEMI>Ra=S>wnPjRW4ZX6t~ACWlRq9-ueha25DD>l&DP6^pu z#b1UCMmhSK?Ap#W=E63AjQ|rv4s;_%n-Aaz64y)L=2tZIPD}}-VR>=&{}Y@+amscC zci(A@05zn7k#}B&NSLu zNFZa7Kx(r&KRh(K;!3nL!;ug#v;WhOvdoU%pu(y>^2XJ1?8eVoj~c4R4)ross$tmQu4=Vdkr;u{wHOlb*ZWlA}N6}5rf=A0NW?KBZ*|6LRxBV8twj8 zhoKj^3aD`{+=Bn90EY90MR^`um~%9WbMc;ikt_IIbOf3js2E%u4I}bSr%@+0I3pyi z$;+LnI?g)XJvCeaQ89TPC@ip>j8BR!NwGRsh^uN9zFEk8^00Po6ExCaCLA6zZ=Sy- zb8_0xbD9Zq1jW%UJ=*FQp5nDkux~u47I!>bN##oc6as!tIza97{cN_o#{?gSFElur z2Fw=F;FkX%Wj-k>ubL1uC|cXd%b?q@y94QGrrnp;-^*xlf<kGtH^Re3N zzCtbfHS{kf*Y3IYvsl|M@92t{Mrd&w9$-}mG38#abKBfxz<5OW<_mRYTjt+hJfq9S z8KS(#o@sb`CHl-->y6Mv|HJdNFy;FWG||$jTL}7}g}%7`ZlmKZy?s$u?g>726K^S? z;5QSmsG`H%`@5^Qt{WZII)ZFK$A?WZi_}lc+Y8aa`NH#N6K#Wf+b0x z%T9<2+gIvn7aJs`1XoBY%jo-6m;dPnNMjownVoQ+(+WONVU8(q%=Cx3Jn@p-Z)Y

&W1bHSGekqT_c+EZAXgF>+;PyO3qwkXHdy@?y9~@60wJY{Dt1T1@EbL ze8?^uIsbuQOnr((1c6RU%F5i2<6fy=NvXE8e_xMI%Rj1CaWQn6Y|BiNQo1a0yu62E z=FE`BNHicjGRH8p6H`O8y)?XG;9d9BGkhxM^y5-waauz9n}gbzbxrKh+~me|;JLA1 zzNGUvYa6+IM@7`R2}E-xui=tozg;0zdv-V-x1!s<(RZu2o*B*zjGicof~3T!7sdQi zBp+4%PR^XxJ($mv0-YUus?EPSefvZDkpRXf>e^`?Rbm1?ZsomfQK9d`Cu3W;C}3}) z*8gaE7eLF0Sd&F}zjLQFE3UqiRn}d_%RPuh%pJZVKTr$vUJ@b!hCM&N3Y-=Yx@()=4>PI zq`Ku$^9o;?!)`nytGDxnKAgjxfB$(;mw;S4X){Pzhc(&CUH=oHEouItHUS%^%z{?- z1uB~lDdAA956PG>>psVJ-sOPwN~Ox1eOVJR_)2Aw z8IL`7EGX%nObu<$UNK0%<+qV=`5hv1ngC?228+i2+e*_1L_dK^xp6L^V$j;ixF0b# zl=!)eqvhZkYD6!d+wHDH~v%C&*A*Ej*mhj%%a`ezR3w~fx4$P+j@ z+v7>+F<(oha~skxI{vPOyDjjYf4#5RDWtLs`#V5;9SE?k0QcRy=4Bx{P_#>8kxvJ> z+#^cnl-y@kd-8YHd(BPn*|G22L7OU27qx(j@}oJLd|<`xexBDKj*YtEv`Cg(2?2RC`-1OkpeSeKk^Z?0t4 zJ?`(5Ihi-NPkHNPdF5Hk>ZlRPwnX-RB@!t-miph;zL1Mv zj}TpGq32aDukuBGCbtC}T?0!TZ3d`zuDyxbO}Z4nGm_wN?m%xWm!T0ODJ^xD%~wJP z+8Wf$)A93oeiqboMdws@^X(aE`zQHcXvjj>Sf}j>gpMq@dr`6F+2V%+UhR&4p~{}$ zN4`@YOY}E$HqcW=*WPg0YwF3#1NW!VH*ywvtd<|am zTkQVz9BPLIg%EtSQSDCT{#yNF^bpDE4eR?}K0jEfLezte4uThk)hKcII5{25 z|4Pa6zW^dow)?sHobe%wX%JqY={EZJ4J8FgS1fKMm({=iB%54Pd-}x83-zqTD3N?e zI!HsD_-2gG>B;JIENf%U!LQ?`MLQ)YsHoN*F@P~OOHSXIQ(L^{!UO#CKqo-? z{lr(Qzt&c9H0M}!`UIeSZ7$-@iG#J0ThlIw<)iT$-xYMzUwx1(VtzplXuQK{M=c904F;zL}FHTm83(t6U^ z8y%D@?>T8LKK>FsPed-AzJ@yrKK&pO+~VYaHqxEg@Vm+Au_aB!)C2anjl}||jXx5b zPRGa9e9xp(>(e&~GnQ|JYe{YNs!VV8StrQDb^qJ6M{ot`cnM6(LEO?~@oa<8Zyw~l zlqms{A_v6%%m&Uaa=|?_8`5~CI?>-jAf#|)ez8eR!Vfmqxnf-2ilz z$eXoFPNmvTqkWsGBr|2Z&^EMkF~2?iy2A8qQGpw3TGrF@6Hw*#+38#FMl3rjWBg#q}X4r_zq4m>4x*4l7fJqJ@5438{;*F*fs#fYZ6??5h0VTKi~IOQeQe#flZ5A?pADT>%{W&Q`y6!$RV{u zgGdbU4m4L*3e@uR^ZmG`q!x~Lc6M&O4!zj%{4W6C2HHnj?{UO9x=TuX)=#OoTH1`d z-c|{dIe2yWgIrpBddCIQZ7Z$G%J@?%PVwm6a;p(L0oyu_YHCs_mt3GPX;qah$F9qJ zyfT*yEGNFI^fn=3XlSppDqI7jtQ=s^XP=yS8>$`1V?lT0vr3+2kf&Y8L=2$bBxhz= z-zL19A1m{Vyw{W&rNAONu&=nZ&{}DCh}}B;{$suCwGs(v4^w{EQ2DrC);4eKV7rW1 zHM787^PbpwR=}<6cX-~8g#|uHxR0@` z!E-noNh+%ru$fBfIOpJ+2?7q7(Gai?VA2>j%Wc!B1?Gl^1Q!hhbmU9jA@RDv10;!4 zc`g$M{z%(#-a0$AtO)R(ssY`4Z?&;Dws)5(F<>Wk%MlHWc%aYZeY`M>kr%nzOiy)C zaBI%|67nI3cp0zw_=$#+(wqwYjT^cxZiav1%g)~GloS=h9{WLHptrrc_47eJ=cTg! zH~rzD8JM9rIz?{PHg`iW-4Iz+avWE=ccr0d9q5aY(OOxxg3?*32s}fpNT0(Fp`!9J zD>{?q_xy&6)ZW7j{gb_>{GvGGv_zOj%9WKzGwR;p@|w;8Pl?xB&8S0oeWEtoZmI%= zSY^LS)FFK=25CX%5{KBb;fhAT2?3ODj802)N}yo>RVqsP7pQ+D{**<-p8cA_DXD*H z^+rv%iG&pF@Z(&ixlV)kx8-8b!a}l2F=_SO^^jx*eNeI7f;P0AFW!{4p<1+Pp{F^K}`*RT&0fz3Z=j zkK~qiSr4yh+JAW%+f>WpF;`vmrQ>5o>LX~`D@fVKr%@RwgVogZWNB&58&9vjo{uPm z!`+)4@vOzf)nU05Kgyk;vvGW9GsZ7TV%sAgr3pE>&63h)9F)drWEa;|_&Lh&mZfu` z8-i^Zu@dk|XIPwdKISf)qvu$!CJKe(vvYK;%3Q@#{Kblt)*oAZ!MBpHFkCjELyf;m zF7P8E{g{D;?Di!b8P{zK_tqXo2HLua;99cBhE!CI zHax9vp~^jo8aH^fa4DfJlGER|?gS?)yv2U=wgv=Z>xa&f zwewv*h zgaZ)^h9y3zBmvVtpm6F!_RAmqFAyS#5dS?NkM(E(h3suGG7srkU2v-eJsF--9UMc& z!QQ;#ueNoiq+){Mc4AVI*DcbnX!O45Wjb0bdrX_C7GrjHyH%PHd+YuDJpI!X3lwI{ zpgKRl!rs3WdxfH67VDSqI1LOhj`2B$r%OIMwgs66)phveHa>Pcdi}08iZ`muaqu+l zl~jvrn7HN00_*-j^)WtSyCFD7i*Qrb3l}P$e^Nbe-8oLnml^Aa+YXR<=40u(cAOcn z?TFD`l+m9NQs3>CCfdDP`g7bg;7YTnn)ru*bK%jw(IOn-8;vcpi-*c~qNPbpy~ob* z;Lb4lJ07l%Y5&SsC^`lqOJ#$Gih}#U`h1=W2=&*V^<9OEMgS*z4i308POYXt2fy5oUDIYv1i5eY}P^c4*kCnEz1yp0j3xmy2t$JE$Dk!SG+O z`u0b7V$c}P%U^^Cpxu?D#5prtD6vo1W`7Ir3$g!nDA#;cT{%Q&9y%yy-AWnCo-n;z zCvQ24AAL+0R~gI1f-FR=_S~*I8>z*`o@QE8eT&w6q+3HFtM&zHS|`12l_q!VapcxS z0_C!YtFN6`jJSw-b)PN&gLVLI-QVA_O-YCEx%;LceOSKfU3#mD_NkMl%LN_(ft$J) z<}T!3ZrtTDI;Y;Pe)!DqDrTtYV1YAJNwK?#P9}}h;1@(+1FmB-UQG^5lOc|WSxUzZ z-gwSAH9NGWH)!0YsB*KKFcX>8UgS=${$P2-&T)Td z?*$p6(xflvSBtxx;`KWY^=5}Lw7?zhM2c*|pRcK7VtYqz2j%|F)PH`tJrwIv08=V| z+y<8-+j~p;+(lU6gXO`hz0sEbXp&|MEBmg)Q6`a;>>BRO_8hHQd^Psi$2YjG@%O6H znbdwTdDHJK)WVHp^m)*p z;3YcU@Nz9Lx)1ZxVE>4B$@_$5^x)F8u(ysiS~q`YPbM;;5$P8k_jv4!%1^3=yL`0MD?g)$witz9?2{Dy6;R_j!p=aK2QiEH}Cwz>leRQ$uE>*YsJ_)IpG-74gr_u!;Awj^+Qvc z!=z825sFgG@Qsq`>TTnNW5*%_ir*S31Z=nD zu-UAVt41d+3beh|tCmFIxPn@lsF@M#)PPf|sHGXFEowbRD2OqL$2aL@V`1O(IhQ6O zJR-SD^4Br9`n|$8_)MuM3EutE_RwbLV?<1<6ioWpyN8xmX;1%_f?>Wt?8lkA^MvkE zFxZPVn4af0s2{Ui;Tk#qkep?RC-QPR%eda=rcEj<;L$Ne0fy@sU`Il8N6ETVWNN>V zSe=B$zI&Pp{|Ij*UV!5fFE=g0$H?>eln4%#85e z3RSFd#{&8oQR5JZC}OzrXm(TgNrF^tle0X4Z3ufpTshVWcOa{rZ2jDPQ2qXt$0A|i z^BUdhK;oN2LaP5uc3Ob!gCWLcuZk$f+{>|ACALpj?#`rYogAlDdfTOMKoNtZkmAOf z+jjONNExe_;q4E2eV+XvS#JRqRoA``9~u!P1W}riMnI&5p`@gdmKKl}1c{+RrIGHI z?k;JR?(UQtdZZh^Grs@d^FHtTW;tu&teJE6zVq7GzR$hM=*FXpehVceT|%9o;`-(5%P#IFrG^_9`e_#G})Cuy+YlbgwHvL_MTfuA~88-YAE- zUDCq}qk;;?*BUeNz}|)}*)BziVb&#^YyTb}LT2{reaNV4-|v;Sf+Z1w;aUD%1;j)> zyfAX$;4`c*`o*hN;(r@ZhUf3+)Py0}GBPp`0r!S9f7a8^wm#vj3Gfl7JNG}v6FW{1 z%*byTdK*obAX2Sso73N zny?6`fH41NtHC#Za4Yoew!-0-%`1?e2Q0ReRTbJ`<8=19NGt z_n_D>&SPyk-hb`W1s}TqDWHD0WC?@991`^2f3nmDuY#`hyz|$+#7^7*Y}X#!f6oIV z^yK6Lc&x)0f^9VK$IEa*3!8 zV{guUyY}uPDU4pWw87vmw1x`a-(_$`pwxJq?qTUUI0-EgV0FA5|6*}dhj=%WQnK$x z<+0Zq8re&~xoVX&kHK1-v!mUVEXyJ@!a6~wpwu=>udt{1)RLX`hA_!Mpgw@(+9O~Ryw_B18D8N zxAU(#A6voauXU6@4X&w@>Gvhp8EI@BXE)C(87ydA^^mV-Z4d1Y-DLUHHRjmOoKd?{ zF+`QS{;c}tR|sYu>&5Xe`IwSVM)hiD*-mR->$}#j+$^kfWkq{5r{w%NuK*jIcdw)i ziYUJ<JKR^RPK@5}@zxhVT^Da_bQ)^nt0W@-Vd z6c|J`zEu7OxARNP6#IaVQ+DlXlxEGOwVo!8`d$3(;?n8?i4D{?%XaSbQvxvulPm@e zwr6BZZW~{inI%`dLgl^%%aBQjEtOu%7Lpg%xN9udSZG>^Nr7@M>lMPHk@=Eat9(tcMjHo z`0I-_PtL>)OAaqfRb5?e(+yGAK#1f*E`koTe&Id$xI#a4MWVBnHIlh_@0q_O`cdD{TxVB z_xK|go!iBGrADo!KJl!uEy2X`rSjY~Z*)H34fcskAr(2CUEB0&kEZpq&3L|CGxw~i zrey2$tJrZS4$v?D^Cf@d8L9fX6FBEnmb5b9 zxY18xW%-j1tg_gRSALjlMh7Qe^I4@=6%u;j)-@7&nG>lOn_?8d<82h7*xOtK zC)J*La>V2=!iR~N(gRVA1IVEk>$xqisD;=!=ad&0ZkdbAmHqxxZ2jPMa1PjH@?Wo; zc0*xE(7WER+NlbzO&DE^E;-B74MWwY?9H3DKeI*RjRRz!Tg?y*Qgjt|qSxy&`7wx+ zWfhp>&Gck(yHPBRK`BW#i%LY-v38o@9*`+0SLgVwl}txFm>D^hh*<188_U`l>Z&Mu zeiGZ*{FXQE<3qdi7-;F5u=;3i1FPcEF79HXWHY%E)T@5I;@U9#y3o_$fj0`G@r*}o zdX$mv<@5Cnu*o|;T`9%!D&ChaLp#4=nlwh^?PNW(#iOZ=$M!f_7v}BI!=F2TqP47; zBQ@XiY!fGi3{Kko4PIi>EZgIt>krrX{Y%3Ed3S%^o|%Vv%bm7NIxb3%r_%AtuYYon zd@s;tfZ>;@(3lH_hB**CP;~LHE9FAHe4u=NDB}C#n!PIj3z6s|yQy zswjkfxXjAJVSGzd6RC(@%hzU(7~Tp7LY;W&Z$p$Bl4u16pBO9d^3XH58LteC5|A^7 zjMzbVIklKLMfG= z7s&WZv4TN^f{S<%kC|;f|BH)7vEig8w>{hBQ$QHD_tK8toDLb<1AL_(a(f9vi1$yfJWx?LU11v<1+F@8rIBm{zxBIX?WKwT9P!R_A~1oCqAv$g96vm+*Y4Z+gQIX% z+h15xIkc#)ZWip=sZv5GC{(uL&Pv=m`tibpigf;zeffu!8I|__kLV&a46~(oy@~ba zygrX}VFbWSs1Dhw@Xb<>H*S65dh(9f@l2ImN!mmiJgQ=w?W0qz#;(%I@$`L~^>;7@ zHN2eo#A6vr1<5#fD6`pZ(3Tiy5yhG@kz{neb&@R2`n9N)W?GYU;ZPe2C5UrfoKK)0 zdz%-G%-M34|3mj>psaw1PcnIYygfVoZ4s<1j)}^u>sRK`l7BBLby1Jt-Y-8pxEIVd znx9@BQTAq}{O&32`IpL1ZwGprP9TonTJ^~$_0>2!PEdTa zhru3+$==?0rf>*wt}n;e!GCn7I{kJqmKSS9k!^CWhib!wF407?5})kG&tsUm3m*0A zXrhYVj@rZ5IE2pHzeJDDzSUXMZ7i&-%=u+4qR&BLJ;0foy6-$uDwQw%@fm<3`s>{v zGKCSLoO21YUwOE&Ss%!I#T2W-#PTv4fkiWz8eUpwDY&uhQ##~)C zLoMfSocUZD{ZFa4^u#f$T<-el&t$Wt;ULFZHC82IrS5Ylp)Iu8z zmf}cvw#RYFIs`;eTp~%x1h|^Kd-+HNUyjfk7b*|e~=PfTdapd#3^tsd&V^)Yw&ogaEg=sMbb=1oOK=kvX zp~u#ixj@a%2oz7{1dFnrtD`ff^zy3y34z#11=umH*anW(q0;M(_A@S(Q|Typ?#kU0 zfu>YEce@|DLJwqop<-Hps*!Z!vDf?Ws2HG#heU`O(k%zBfFaLr=B0tbb0Rr$tPYl-7DmdTQRyo$9t=Zkp= zo6h2llkE^9btZh}^it#0<$F7?RqZN#R_))qBgA`voZsmDcUUV8wn3?};@B0(_k$Z! z-SW)N-Cq4$g@o$0l_!3iKLRL?Eaz@)V9O+1^8HC7%o0?t`50Hf7lf|Y+P&tw2skl zqDKHJ{F@~|CxOsiW_HgDuTQil!#~wK-W9c6_tfHnt#V)a-vpSS=X<8LWY)V|e&?a} zxhbC#d8Q+5rLw)t?BdIrllSAdxvfj%pytw?<1*OaPSNSub>nQ6-Ocx>b2|8BH-*su zxUYB!lCa?%9qk{ZYo}|1`h5Y6bSV{bCrUiftF5TFLKG>R?2m)`JZ9uH2UIdv`(44K z;`-7TP&n8qNa7sSI3Jr?EF7teDdXm}kK?KuS^Bvg`@9m09a=AwQI{20tC9(P;la6U zbbCD%pwe`areh`I;z4`)*hvky(>PA#q=uH(SU%LxLV~!pe1V!oE+2DVQe8ITbQY^Y zv+?w&4e+NUQT}z8X>PRw8^vlGUmcBMXV&}(l19$`w78*Yb=07PuhG(wDUokwlg3UT zdwf)qmamd*wzFN4((4lVZ=*kMMZ%xQqhaf=e7I3>VN(Z^)jPL-iH=_O36WBJ1Ie&Z zE~}39Rb!RPHdixVGkd+1ZAYx+@Vp-#G8dxbkf(6cpUG-fvl)P}iv_DqSW1 zntm|awM2C92|vZ!Qydf%D+Al(qc_H-U0es!9!Kp@PX^xBx^MuZ*bYhGI(hBgzp+gu z4?_jGCgKQ4*_vI={KweBYMq|U6^A+(W|}2OzaVdI#YFLSAF*_thh#h8ZT=s_{~x^x z{0&M3CKxm)p^r6C0}f^dHh4UbA!H)1%4G7K{;D}_Ha*&(w(_z9?Q^gtIX4$x5>%5@ z7o&xfnzNZ5WUq3Eq#MGW-a>`lR!UrJ7iQsu+Cd}EyA@qaD*ak@b=uX0%C;&eS5OwY z9NcYAI-&)AgN(UtQU8h0nRsmCF=r+ZXxBOFi%wtrum1cW@%oV?UCu~4FW9@(`tV_o z&(dn*yxqH_wNSJq6zn8s{-$RhlR>H9Itrz_!q}Ncp=BJSh%+q~CwaUohHoS48{f)T zt8FrVk7?KMiEfd5QK97z?(Mu({Ze;a(Y4qcs_=wIwvZRHt<$!sB%4(*IN6ZNLFD`z zyX~x@DAu(}=ijOhP$WSP27*odtw&~q(-pQU$KXp}(WtU(B3nTx zCufM?SYCIq9kuCxOOfN*+|{xsDW-N`a!noz;C}(}-ac}b(|0ld_l(N(Ui|Nkt;@gDshN>)y%X}+h=1ih($ zlG=@NvH}nCZSL&FLnC1(Z>gbvI6kIpAt|R4$+RYml^N94BOiqV`* zyzQ4~SwhQR1XF2DLna*X8@c*yMDE|)a{3b^pMIG^Nn875mu#&&nWJjRtgB}*(Nl2nc z631UcZ!)UO#_0@ zoVIg7peCXyhu8`KiS%S_F=usYn|?k~%q_<- zgoU6$bM65X=m$>#HN)pr!HS9>Lx)nvlei6?99Delv_Oz=?3;l*V_@5Mb@vkYoy+&? z#_YHTziq*+rqwHvdAH4huGexCE$8pM$TBl5hdef2lrFQh`l`&IW_)J#ShbJ~r*F~h z*fF?P692l3QMP&{0syo=w49vfV5TOpq@w9Wd7&gKgY}+>ZiR*bi{^gNGQwC-3?XbN82sT4@!}0s>3J`W{b3go}tPq`rhQ=i5%cG!|3o%}>s9L^OOxFiN~0 zygu4UHG*z+kErO_|LItA{80~mbPPX_wp*<=IPZ!xePX=o*f7pTcmhVVS#3{Sqh4+>C>K3cj2LC~MK*(sr*I!~8izviIksxK5(Ojn>jXGQc z!y+V9{z2r;iToB@N0RHY4K1aNCqmp~#i{_kXG$8=2bo_8uQA0rI4YsZ63I*J$wneh z*r&$+!`D59j5g}u1{myD?8fqvQ_ndcz84#8kN26sXW{qQzKT@y|2=p9#=P0*DShw1CD#K(WO@M9xY<3;i-yDzN_z> z4+VDz>A z$gK`-$4d}HM56c^-q3{PY03neuSKf3qKt1Yy!x6^Lp2*UY-KUg@bz+dY-GO7>!b-P z>~h|@S+c#7jT*qemEu#Tc((x)sXy(o9s{5BS;O<9C&Y{0KD!VIjh4S(#qc?-zM%mM zCW${;(h~FdlIj?x8%iw>*^87}#Tgl_4+V6KIZ)o_!Q`?cA&0^Sz>nPTn~~Jmv_6Zp zS(LO?Zw`YsZcxKv+Y2wo=cTGO5tuml*qhM&qNDBIM{;;T84UGhJ|v8Yh3s4qza8JQ zarSi@lk4wYa9tEHWNZ)V6|{eR6kr}O;ZeW!J$f4~lM-82FLq-plAH}5 z1Hy5b`ACV2U0=r-1w-CVK?D54h@Xr0-UCVqMPnC*BIH>xN#DP3Qp6LDAn@u-bLt7X zkSyfI33$}>)!bR8n{zeD_6Z&2@UVWAk|c0whg%fYRuX*7V;O98p8uj{b2N_DRzmvx zb znxN)?c0*#M+>F=Pwg)(_9*)`=daG>u6Sr*%!(e}H$Qoj@lUuk;2{ zZl3qP3o6lf?Lv$F!o{x5tLq`b?JhehO!)yu$M^1SsRgbtkNS7MU+k4v7_^`7`-GE9 zY$UzD`{ng+&v`=s@zbLduY3I%&ovJzEhPem;X;7-w=+94;VUm~TYS)Ou&{JWdt{AN3GQ$1RJDC721rJ?{R%w% zy1oR2Zl!I>;KKC=u}a&5Jz;R!8ZW4f^To(alW`NH-Kus8nytUCCt}iHFjKF z5P?8mfZ6pmf$}?s^r{E&ReNNyj~`gC5E=d6=J$FjL{rJlSbiW@^G)SCi3B6v)9iPc zD)&tx)wT>qM6b6}NSm{K)~C}myPRyy_XX=Y3{I|0ADnagY_h*~9UP<=x!fZ%cGx{M z%%g*ZsCoX_2eNdD%@sud5H)bqpl0~Cn8--nl!H-bu`xqk4|SRJRDw1zE->ezT7c=# zQpDeZvJ4P5WpfeY`I-Vz$DuSC-tViHo)?in`CIko7p;;*mtgRj0JgOt);6?PmXW=t zS{3b%0|Y3MheUr#kfwrS^cd%Q$F1!pEx{x7Vy?m_5uOcw8eU8W1drN#WAH`!qcS~XJqm@A&H%D&L3Ldr}G+14!C@s?3|tmnryms zw`tQuCJF4BUs_=MfsCJo;cV^#^chjh9BPEV&t{Z0YC5@+&6bfw0fORhJPdkONf@EWcY;u;qU!S5h z>x)S+s+Ph4LJZnyzu&?bNQ9W-FuO;>EE70KQ>KP^w4Ii$6?~TFh*;7<6}6vfo~@Wl zwG`5C5qBt6y%7Wl!++5; zAH6lwJ6Nna%NkD-l{Ti0s9^Dl1TljTu|NIVg)59;WJ=~UP7^s3_c;a&@gS1W!3zRd z&K+bWJfi#+Z*pBExRcS>(tiv8XNwFwU`ejK_IUrh=fShv-)MKUlv2#d3I?i?GK*$MAyT}XFR_)o9Nx97Zv z@X}Lq_B?qWZ-nJGYY5+S{GQu{5?W)-csr*Huc)1rO)G>4pZ_$-k;is(H5~al!#AHG zXnMM`1t_|_7J=Yc`!BliPe01hD4aM5&EX^_d>^)sy%Kc|;JVflD8GDsOZqEdQ|(|i zO7Ty^wZA;OB*yp61+Y+nxJ0l(h0tV~Lxf1#7U#Whuc2^8cor&HGXNpWV^7+%-XA2T zQn{b&{wOW=rrL`!ok@287nkWeyC51c>{Rt*<4+O?56}LEb96%&=RNWBFXde`Ati}M z4im&))@&?;-{P|z+8NAf82Hi3d!7U+%KLKiq9}9Td)`S$!@^6|8GoMMdqFw@_pPlZ zUre?6;a2~9ul(m;S*wYa_(x50-5RGKMb>}nQVg1``U&W}H=U1~!CI6mQaz_t_pi*R zGXz5HguZ$DI+~gv&2-A-cAIIl9W2Ns2-9nvf>~S4)hWTFvsL|P0!%1-vYVrzt~UJ0)WB6)0(GVha=+OM8$L%pnuL%Mmy?IbZR)*}NY5|&5b_ih48QjlIQ0>) zAQfCUzblHrMIz%96HPqy%m*&m^LeoQ7>|;YooK3Udy_veaWmfeUtf2dD|{ir=r#A3wZYOuremJnw_`~LRoj5KFLGdUJt-=SLp!Y_ zE%!`Vvl^ukxXAiF2KpKt6y%iOpiZ;!_rE5#+-3=0GN;&&`Ww;XjLtc;?w$esuDQcg z1C85GLvI8XWq+=H!N7>R|HvKQLZhy}R{+&@{BYLArr;(L6MI=8pVm+dxHn(D;9=q2 zm!2H2v-UMBei5cu_9^HmF_$KS2{Ibjl$}1ZS2loMaUB#LDbkp&nj<>H{u4$E*`}sl z<7Gz1I@ievG)VCpS^Vvx_SenH= zqo-eMRaw2){mNLW@4z9gfe-)b+Nhowww3&LUO_*zJ%ur8GF2KEZ03BkeT^B&*<6s4 zNMCKe?^P@o{u_4B7ilv2Y-w7FeAZr2n$Z&*+Kgewcjs@#nsmSl$xA-w^)6SwH|F`m z+?9}zv!5RUTK~^IW*v5J8&B!(dlQ2=pNae4XST!kk20)CMZxKWCs(-t+K*}F3~x4F z0^%zH{`P{{U2N@Uk$BL^50zm1Cz-~p_8Xz0FU}k#} zZ{xC`cEIhzP~JB_(rXh{|E7$0V39BYZkp?9jcrJp4ZMwy5wH=X@`=xdQ4sn*XExaz z#6Lu3NAWMx*zpVrX8I-V(vk>Tk2Syt-dZAk$e7pj{qE^8fIU^(7YOCfy!39d{SW!bPIi6L3Dt<^{|P!`g4EEOXIgks%TQ)PG&Y zK4Sd4k0S_D#xNO~d#V`iJm8lB(G=L&!ow5XH@H&9mcOHajsIfyRpq!HLPVv+o<6YB zGXdF8LmkZ%+R_u};t`XbZ%$AO)3mu5)cl-RX>cVGT-n@_cNJ2M*>4B5_?wtr9!6+C zXXa)5Xe9e---pw0Gu{s9+!~gWFx)-=nX+CmqyqvV=^nicb7^VXC1MrOO|b}loBb%Q zT)`#94FTxm$PxQE%av%YPC>tlg{gI5($azE38FU14*>EnrkAvQqBx)r`J#RvvXoo0 z8Cg;Lt3Af-Xt98#sD}gm1aEU-L(-(yitIdQ>;2;mPvd_?A1m82Z)Hji=~J8o%qCzz zp@y+hqrENjJv+y_d7GpT`2POk*=$4h+$;*LR9reuDbG{Hy-DcKwdCZ=J?N*oTFbc? zJ&P=lR96A;E=WmD;hyOeYsE(+fikl4Mv0uf^lSp0R^Rw}=~8(Ld8y9%ecC9HY z@wESzlzI=Rf-d`cW03`}WD)w48&7;eMZgZJt;gY?(9l8ZgBVB6@$2F};=~NW{efol zC_K@~MK;g{f3u({C!M?A1AQ>!Egb6>5CtW~y&_|NN_GzvhPgzzN`2?IC-x_1Z_CC7N42sSPG}c9K9tO< zsNe~0w1_Jr5pxGCMIksI2fvOJ(NbG*XkH;6We}E(Z5}k`!q9zldV68>7m{&}=m7{2 za3XVWZ|~&G!(es~yzr&TziCH(#;@QyLwydgc`Kg9H9_+yLxWMZN`R0sViD{joIJv5 zM~XI`0qc{dN13Wp=NlMCG7E|?FMpC>Wku(Io+M`Gxytd8!v1V_qDYhR=M4xWX}zbw zz(cl$-oA|-fGh0HCK0OACWz-Gi1*daG}YsfY! zMlY8uA3=yi5KHP2_^t~^8yr1-`|%y%8Hl|?-lK|* z4dDE}X=wDQb-AM#x!630=tP4VUBahZp6e?H2o528Lw4WDq9AMcVba3FLP{d0a(k4D zEcUfJ@?!WCCRVj}qCV0}D_;wOs!d!q-)uf4@S1G+4BmBKhW0iCIn5HJr}(z7BJcSQ zFP}4kyh@N3X$I{TZ$E@dGr?uPOCn+ke{ucZ?aSSw1gjtziCeDUGXtuvl0PXwSU5Xp zfGp7rlD_iAlI--LW+@&3Ms~4;(Jw_)f5D@qa8nuNp|>lC|~`41v)(Gn>%DYAXI)i*Kj7o z;j^|6QfO+c>D4KyD^I{2EA?G5_vcCSB;z`=2+`T8p=rSgEXWYkzvm6@z>WGvLVEhw zEEwNfKA3WLOHoB-m%WILRvBA*o)IC*bAE^OT@E=Rk52<+EdcEb&y$(F-_S`cRi^K}jk*=dmT7 zWB{ii%f|hwTgkF_@@N}_o5sY%N)YwX4EzQN>MD`iM!C3;gclAIu}upQ49`*+)`&Fi zXQ@*DON4835|TYI7zNGAu-tJ8JALqT>zwPE1WpIB-orFu0c28ypk{5K@$JnCGk1Rq z>{xl3XF(7meehg}K5FRa2U59NIUm}jMBTstM-_#vOXmeX z^EpV9VgFrr^F%ZS`Q6WuH4Y`*xgA>@Aw^jU;_YXTUYB=_dZa@-AOQ8w!o&R+3<@8I zk=CmS){rft|kY`{iH|2pFP>nP#ps-MmZ+8OvbPeCy%fpa$)D_tO(?Wev$Q8z)yJHaj46Eyu+X!XpX2P(P(U0;} zZc3uh3SBXf*e#E`zP@2;yTU1rA-(@=t4Qlo5E-_?V#Y`LN$8h=03dE*)a0X0Jf*Gj zvXFn<2uN6j*b?bKfD@B1)Nd@zt}~eR3MyhcpLffSxO#uYl@@ zYh~3&e&xTm@ZkCgkqu;^70WOi48Qcpt~U8t3q;2G(qwq%7f$~nrV$3k!ZeD#95P#h zR|?7{;+4(J%0rcT+}3Q^eI(~yP~_f)7np?N{1<#VtaW!jK$%ZR%0u9tLml-f6iNO{ zk86EcfazGDju0KPpw!*-wwWvPs{jE8B^X@^NA7_SHx1pZ7`Le8=LIR<<2vw$ zg}wmCfxHAGeReeF*Yii-k3g^n@kJ*S@?7U_f-+T%M;Eijwkm1`%{O^*d7pyjE3o@1Is7_lKzlU#7;u3CC>m)QN@_)?RP^pX!ZXn;N# z+nID_mZ=@rcW%L50?C>-f4=GTJkleqU=YdVmNc)i1s6OXiW2n)p$&23b`akr4l_H} zGDy=g$?rX{U#Q)pJTUn?EK@zbP@>=trXFceIq_m%ai)f29FcZ{0)lgy7VY% z|3QESB_CEzh0}7bSHB_`i zC*X@ZB8x+sIlE)jzSU^!HmYORqpo98NQ(2kQ?|R1X8lg=9TZ4fO!*oiiqe4P@66zz>$g06QeS_Q1KVIY#zE`CMTPT3 z7jKP1eWX%ngB*8AXl*dRBbP5AG(2~g1^G19n%}j+EvOOY$D6SQg0(g8R}GuQQISs( zNT#y-qSOZU*(FWmpl#C}xB(A=MqkEkO+;g}gOdPxk;WBS@W8EcaJE~1$81W`V|Qb9 z=E0cOAwZDPdx6M&q~XL~7@++5h%0w!Xt)eq3ZYJ)&EFNG(V!OJL*poyOPA@rnsAP9 z#(S57w6{y-LYrGO6%Rx;k1_}}a#=dZ3T3DGHYvrPYAS5idx@@GEd!{8gd@y_%z9_| zDD#-DrC`kS&8a`cYgvm?3b=e9g$kcg|1e()lYTDMgccCs*G~yS@zf0_zD94&0xKDd zri3m32!K1C02%P%@^DWqAfCIDe<+v97{xJ}0HP_I#VVTUBaar28K5SbgBn2H%IpE` z0?*R!YURXu?+5bIN)!^Wy!<|i)+tapF&iBj?_5oMNR2|(m6XWmGN0cxAtAv{43?wM zy=^oG%V1C_%HQtJRog(eOTj2XWw#!qxu))%ifNms0vJ3dJ86SyPIgeALmcnNO>O~9 zRnpA#-VNXm5)ayW6Y>%{_le+BF~v*8Flhs%XZ%KB5{A76C9nWQexaDa`wZhLB6Cua zCG~DgyW5JZs;ZJ`S^D@!U0s=&{vP-|Zp*^E=jTA{hnhMvviAs40hP%9i_g67D&Le9 zf=1%`!n?q9m#%*Dd-mo};eE8?{8p|sFI6StinT-@Pg<&T#l@WoZJVmRE#ub&^u+F{tn6{c00(eP$+YupP4|(7v?*KHEkd3Ob@q zH-~mLPOB9FR`w7mZ|1Y<-1%Q;$SrtZKXjf@gaGE@V*I%!ANFXQ#1;Lh%Or@KsR*Rz z*=;1Bv7lUI{=@aoUnFN&t6iPdc{95w+D|$gE~b4Z-W|Z{RTswj*ZM0(P>K+)^mQMXeicI1r96D{7_Bz)qAb%(J>17 zI!(@pVJ1#r>2ij_-Ql{%4zIwWq;Bm(*)*Ma*;%{DBuXc+pVOy_ppK!KN8GLk zk=#7H&wvWmisAr_wXj)Eb+g)!_5L3I2sj2MI`GHKRK)%2h0KTpkJ-*}j7ecUdM^e&iU2Np{Cnyxvet+|wP89}c6JjhB}XP3q;i-+_!XOJA(Q zqP8%NU41e+D>I1`6)IHYqkxO0_TTF+BnP@wrm0sTDp zLr?o*PmB3JHwQqau<*ZT|#ICRw9 zefzIsvwF+G;}yb#KlDbw*Xn5}Uz^vYE4~lK_uj2^>aiy&BvlaB8T4e%r`;74h3f1H zBHAk=1NqigIH03o|C1}d&H94w8&zn|0iMn8xgGWIjI*>PbSFDg)-aQEU6q&Sh*Xp= z6ku14_oTx`pC#tjK4(m1;M>x$q~g^%R`Y1WiYqqIP-gVoKBB@A~eDV+Jox1Uq`jzXVzqp@Zxp0NIN&Ws2(E6a7M5jjWqqdtveWDCUmr?n| zD>KyDv3mZf+PIiU1`q9=in6v0;V9t?!j?pV&H{QWGN62Tr`lyy|7Tc$dF-#+QXBoAiIJ<#awSkM`;|sGzC4}lfgn9 z9P^{R*yFBQ^SxW7{A=$$R58Rto$X5;y%&RYL#7o<_AQr#Ht>K^?V>6NF>$S)ENJFP zH%jCJW8xZ$dnYm+tY>$Kgfe}K%!0KIqE(4zCU^3iUB$-8HyRG~->WvnIt0%<1}cUg zpIOU<$+4KXvCHvD<6lH$KY1%*!FVK@;Zpv+5&-BDW(q7gaOHF&c{^7pIM9=^v7^C{ zBp)&s={u}_Ruq_%v2Pnhr*Ydm$2$%@J9+c-Ipd9lu88lh1}(gkLH&6VXjPXX_YHoN z(K(QrAe8jAK?AC>cWNIDq~$J01%|~~9|_mi6*-I!C0C};a$dxiGNQQm_Kw0$!NBar zpV1x2uxfU%+w>$dB&>U`v?$BDzaUa3>DPlbJ}`v5kAuKUs6+7M$_Jy zi?ul5KXtjF%-Nv?KFM4_O9#&)L)^@m?)v)4(D2ZG+IL+)s#_c!nFKROu>&8kXl+JA;yh!V5=nu&WmEQhX$B9QXFCh5ifg$ zAQo=Z34qKC3t#>lVX_chHr5oistSige?v1%4i`hvjlxLkK1yGP&CYh4UDkHa^fkX1 z8x!1NKfCA}x&Csw+y(#}LZ38q!de(4e;ARbgeD|p-N{}e{tzxrx79ULyKG4DZ<#YM zXmSDqTuCqJnYS!tlo0lRV5hus*Itk8;QV6&WtHh;T2IF$qHYQstDHTj4ZS~9TK>*5 zi7c7-x+^t^+Ks$pi}GAvYlh@oP0xUFh@=m$Fb&yafX@%OQGYsRR*XoKc8Yz0YX=Np zM6;y|%8eFzggtTe)m%FF=wT1va+VnZ@1dF>e%zSs_8#=Ts+6T8#2=)uQ+EDntImul zWz*b|9hGvaAN7|o$FwhL7=w!#Y&dSs5{OX%o8y)bV`2hN)?a6r&h8vIyPK2+dhe=e zT&_sT)XLU7-nF%GWpTy`=u~i}%u64HSORI~mIsxE^dT|aPMhgVvwTXYY+PsRQIznk zR71G9{}_!$Xit*`A&J*p6~WRJY_&AsLf z0Bb%e1~A14?v4~!%SGWN93>gad45Hwf~>C{9O9@ci@o5X-I4oI4okPML&%1wr=4rK zJoF9~qvyTLj?&v?WLP4&B?lGQMVP$m72?BjplZX&W_Qhy+N3O*NY^Sm#wdq;D(;fx zdjXZlQ5T2xj&L$4fIU6`+J<8!lv>*(R0bcN0;}7yk5srpW3(x9-*K=m9V^YzcDybk z5~@2&`k{-j@C=ua-uuIWVNZKZq)n7V67*!e(WB;6_*WFUMGhh1)^Z9p*QFoD!@5lE zPjd7+YsqY}aMPL7_T7pKW=BITv`<%`Z7a`k@JLuGqjOV>Cwdu zS`@H$#z~RY%L7hL>AXDCz_YgHoZB7uEaR)mK0(hZp=zDCq7__i$lvL9q)fcUBmDhS z69EVW+MOj%w?}9xy%s4TBS+8XSRu*42Q6*>or07(4)pCbQWnBvUq8*x`Fq@!9BAx{ zL%5^tOOmpa`Zb3;?$bU19P~|VKK0&D5RS7^q_bqeNW|)Goi3c8ia)*3$j&Btl}@RZ5^~|6;eft(yIJIbG)lfS8cJ`a8@NU`B zRe>Dr7G)U6{pg9PGrPT8_tqn7Cv6TTxK*-UwN-1f)@6<@X*PHuBP#b@Ih0GU&y?97 zALCzrv_%V%-3kR|7FhC?$YCwC*R$)Y*pktv6?radc$DvjA5_nyvr^Wj&(cdo=E z1k{Qv`kBdoItb(bd6_?CU_`;54Aa9;XrA(;jGPrmoX6&J*$(D+>rJw|CDmlR?ecK{ zxxfHswX8ELDlGR3GRHI*k^&`+T*?CfxeQQ+G4KfZGNBHa-3BLtQ>~GE9-U+VxVTiV zUQ*JAZ|N-bEo8x%W2SXY6EHVBqH3y?P>FNki8>QS!k_r2YWj4Ggva!d<@4g%U_gd#)z-zq{n3!6Sk26ntWL0&QAEbN6 zFcs=ZsJHT9)EpMr`NQdl4ex!u%}eRagWOd9n_;->$`)0kR(hN)gNN{afF|E{b@qfO zX^Ju`T`;?#VEZe$aQx>Ee|$@)*ic*9&*%DU+n)0a(QUDZIB~64$&Q-JI7VLfV6>M^H-D z=W@uo_JtV>a#Rd5d+-?y1Zcj0#dHOrNcnNaYyreiXZT%luf;Xf8gS< z&a8@B_-LwY_u!w4tucZAiJmES#iS#=@s`ni!@ge7c)%gZpetT0B`e6Gx))z_^1|@} z@Xmh}PPLtuuC^3@OohY!-*f8#1C!Urc;2By7I$S&-Nm|^P0v*L#4IZe5w`qrcR@S7 zdVh85lH&txp!M>COOD*CmO8q|+_fMl$5gd7;-7&4=K)A*^*F3sJr+29-y97tB=b$1 z#PoJoW4dw<9l{EEOSYyJq5y(2E&S_h7S(~M2FC7Z|NZ?nc(k|oyO$^TpuZR$=1q|> zA-TDw&r#9;+%d~uJ9hu#bUGHrGX(`mS*T&#;iH=@x-ffNb0n+EV;? zG<(P<$UmQ-;2=lSZ`xeWt~8oGpxcyW&?Tj!otPd?ZO~*CVgLt3%_hp4>P;7}DnsRU z9A*g$3tSnNrT=@OS&)PPz7X~lm^Pz^d;%fDxy{$+r#x;esUHX2W#BJ(jx5GTFOqHq zumFL=nY+t5-GuBE+Z}mWl9tDC!v9y^cmFk=Jpaez6g@<2h=711MY<9K(gGHgP(6B& zO6W~`FX{;b(h`cHOO@V1N`i{?5+u@l4WUH{5IUb%yi5;n!LhbBos2%5=6wlWq$tKfla zmC@CX&rkgEYnpC*s@rio8!>PW${l~a)5lMZw(bj?Q1)~J2Fi)6D&v+MeBpwV5Hg%F z#9sxuJW%G5J4lXGqE|C?+3fBLwEi^{hz&&DSkGm*0CmS}667$q>%F=tpJ9L?Td}Sh--5@ z(F``;eVMW$PI3=sy(s*iN#60d5LZh0=>PE@M4qB|3iDJy8b3l@rI+>Y|xYG#s|@SI(OdDamaJ#~47gJ4chJ!@bbmI= zDm06kG5AYDH63U?D(*q zsGVsuGyJZI)61y-nUKv*L8sj|Vt+v!!k{mg5Oz3=vB3m16)&y{D%e zH!jCGEq-Ngqnln`AOI2MFCJe#`-5UaVuyQRxu zgW+dU+Yf55&k!MyLzWSTya<#2%tE4#mAlxQEI%x3JylUF zf-Q`vQ4-anvS0B#`FCg9Q^a=t^m~SpT@J3|_`y$8wjBAMJkprrd+c`>TXw9Arc5&*O3Fv3<p2$>tH12zRG?*T=D~^>qLN2 zmDTEACSWmr`HVX=xP$L^;(>h7e;Z5X<8fCUUZxUXhr|#Y zr@Dh8%S40yapZ%xb(-b3U-~y6K>4!nEQ8VN$Te0?5k)2G=MJmnh%C$tCQxxhbiGN5 zF>}_M0?>vb$LGH#0&Gji-w{^0iHXI|Tdm|PRbj_YGEy<2WICi5VmI)vtLO%Dh!8)j zJaWg$>s}9`sJWB9v1SapT(KQ%H!Q9gpHx_BzcT!CXweAN;+mILFH;K0uJ(Ir>$gy1j<+FN`VgJNpRr8DI#uE!rceNv6Mz8wbAO>8@FNzb>&ss&EiNdpC zxmS1NN58HFBn>Jtm`=)%rWP|$h!p6eack!j&&^3(ZVRSWq#-$(w)6g|7M<7KqgDkI6 zwY0fhYJc{5N|}n7#7p#$r=+|Ro|TTn1VKw!VFfFxC;qL)#GKK??!Wk01a$#K=J z_?R*O1*|r&u`e|X_5D;R&%t66VsvvF96cVPt1ZiJdG$XgCxtJ5B`mzQ@INh0TMAdw z0Q=7+BqUNDYb~nYhTE*gkDh>`5G+Ns>rGuYo`(_H8&Lnh=heaEjrMX18(F3GQR>K* zuYW36P_C8eb$-VBdike{X|% z#6B!L4jC6@(9%eER%iFfZnCgY21^m2N6DDs#&E|@14WSLllBF_$uMYtYf})-eLbxk!ziUAouk4 ziTDsEv`kI!MS6m+vZ`XXb;+IiaNy|g^8%ezPDufYuzyFAzC(L9u*iEV5GU^1CN&Kf z>}1?21=zk%@6%*+ecpU)iocQ?YTMb|nj($J*y68C!&+LD5OkmYxd z*W7yZQj<@XGY4vHH|gKR2&7d&r9o)STD<}Uxd#453T>XwdqMxxa>!>_i2=fv2c%$B z`l%{O#n`a?K(s23GoI55J}&y_CztA#(%hTjg}`&mg&Cggyxt~(@Ha)q+U(D`KrMff zJOf{-J#eqQkiO3jM zWzq0vf)4_{*NSmJOwpwUV&PKx{mO+ez?c7_H#J9F{dBdpDkS6%uKatG!|6=A zPVx!F4O}p{HT%gU?;(;4T{j4T42U}G8G)7`WEGfk4V%PzF?)Q$#E{YL zG#HHPk;r(ohFA)1cQ6^1uH9i)`7{zra*+UHj9TR*3db6Is2g~J)E z;khlNSzlL-6|{tj5=gHtj|asn^oOlV#AsfmxXy|Q@#KI~}Ac>(G+ zu$fcy_8;A=Qg8e4aRg$Rmm7QH@~NLw`UHYy&jIkeb&@JjRqcUeKHsq_n!ZCe$!o6I z!m7l#))vB|I?P#QA&;qK42phZuZb?;%r$*@?Q*A_@!#*?z#VEOrnI%)X<2@KMBhgz zt>mWMviy0_Ntz$i7OKw~zI&|}hv(tgG!P;t)HGZB3<#W+&L(3b3+FPC_sqj~!T4>E z%SjPq^5L9kYZC?>Jso8^{mQ6_?mtrH2bSb)%+=o0OCJaMP$TE+EU6`sU9o~1yx+G{ zKdaMj#0HfC7V8q26!38^8BQf;h?H^ zzt$_ZR7RE=m*l457ZEAqm5^wHwSx%rR?^h<3e($us&mfN22eD%kAul~7a7hx;03)u8P2RJcq|i~h1@EguIfnx35l=~Dy! z>P9QWfpZ2E2BQuXi+nGya`bL~j7#ZC*)}$#?fj~T%p9WMEEng$RQPfs*VHOuA9OV# zKSi$0g{Z6^)}bXqOwvo+q1-Y5>92y&1LfM>D3jq0DbPuFs#1mm_s}Oh*O|h~65@?k z15BdJ5lp+$Fe!y;+Z_0(QkdZ1V6~n^7YKNfvTzu#e-SV}mQh0?LWU;QSj&Kn7cQud|RtLb0JP*~_K{lSy4!n9OrY7gGxv(mJrzaXGD}l+5R{Q2#+0|^T&j+Zi zYtKvc0)Ay{)MTSK*dr1RdK}lW^eXQ`r``J;<|wORRAFedjR}J&P!sX*Xqj(d=ieHq zE}{vz7mITffZ;xBcwZmn?$ z^V0*lRxzsuq5AkshM1|XH*_i})8K^!nxhjZ3T_Pn?gYR61C&g}Q`t3gzO~03nyq^g zzYVzYLVlEEh62Q9B+|+2j;u@IN-6(1*u2C;hF$Z1oO!Npp^?*P+9J7K6JRe_Q3@z_ z&urT!0bUEE#|9T$+Yh5E%Bz1uKdAgdZ5eO}+WsUP6kts4Lt5C11y+(U+o|NyyO&3~ zY$4Ky??>!Z+E*EH;Ah|HliOwKpC+1i! zFDHVImX=FymhDfLl@6pKNW%4Wfc(X%-Q}E=D~-!+ybS+NN_VWpbb8A^cCb!_>_qZg zWNNJZg%Z!oL+hus>V^+loA5Kb2!kRs>P_z~)v+xnRSVEufi=Wi`-ZSC9HoG`1REZg_6E8@$2 zdD*mA=s@3!3u+dsZ;j`Ric~+$f~akVderT7UUmK%83(+Pao%jursEZ&lRgO&qb9nX zLgOGDRbzi^z_X)p{v~?72vn_zCWgxCSj_&y-c|6i`4J6z*9x4`UPUS5-smxACP{a7 zv$~xFkm;{^v-WH15Iv!QUlI`85}S1yWSRWhXy6=Ir>53;_dOjpx5)5Z-J)E$GJofz zaJAa)+;_dK(7l_0VU^+qOuMW*Q9pA765IeTa4glGTFr%@K>xH0AhAy_xY15}?$7Ri=n5(KQQFhQ6=U*$V4X=`&umnk60DDjB zGSz98dG0eM=s_fkQinawV^0hI!dA-{-WfitTz;oGd*rHKl#ye z=1^Gf_8}*@toS8J-iCSuQ}v-6YTv1&1vlk{7f-D&I3f>#IrYlcwun`5X~>8BAgR!* zgU#WbMW!yPoM3K7K6MNG!kod&dN+RRMr0IfMF{RP^Gf{5U4EdHsoUU(c9`5iZJXc7 z!CTJIJ;a&x%Ej55WbK6Tq11(wws2?wBD+gE{bn|a0@>JSPIHWvvQJkM0R3ryq5RWd z2Z`R4lyC#C^v%rsy`Kaf3#g=)XV{5Fz|$GYc@aY<|iV;|gn14=+MA3W7Z zWPLwrs;bSPVEtN6c_g}KY7RM*iax!8{Am+=oR{u`O6@t zn3XcmaXyTRxioI;hG7DV$}Aav*pksxolN<8t)+kDY~b#SZh#S4+>I2E6kEq&ntHEMFbx`|{C zMY&`)4MU)$P&eS&=~QcU(n^LIx-Wxek>WjMJyoHagUt3?b zPK(09s=l1p)?E$zd0;_*^iLod6p}UtsI??MWB&*N6+;23)Pviv49i4r#XTcZO<0l7&15Q-m!Y!`=@WU+|YD(AQh1XgXcX}n7O zxbh}to*=coL49KhAR)>10Rbw7%{o{~NZi#HG})kheDRZc-I8e3 z9D0x(GUVSjmQp9SBoB)GQIlrlW;4Az8?tbE)TZ4B_A>t*u&v#7qI0~ab;wMhxnAFN z!Elk})8W1i@R?ovucuU(3Bb(?tMZtkIdb{d)LWXg?c>Hn*Te(dbTChrYY2XRZkYZ@ z{+1WoHF#YfnZephqG^`42B1DN@-G9 z;FsMV0ktsNDNtWE$2F6fa(JNGHgI>{e~MLkPEC40a98bFl;t|A5pmCV^hNWEUy(~s zI_=xIE7xFAb$h+zd&3vP_9>FL@yVdl?i%1a%cK@-sSb8{mIdMIm9&t1-JF|3C;H1R zVGDyn)C_nBAHdo(TFw*7_0}!{C1gVv-JqbB<6Rv<4hht$J5%`{Lhq~QkSjmlW2`gl ze*=+_EbD!y1LrIPTX~<==ej8OxT|W#)9+AN){D_A2eG!F&5>B~J!d!2GiN3eSJ=7O zv;2}n7EB1(RHpx^PRIMk-+L_Go&9=FPEPzONon3odC8ZwghBEvtk9`MZ*B$m>YKbD z?0_OG-B(FJ@Zv1?39Y|{dOP{Vt0bCg?kT=mg*5G`rs)3e>lL1(iLJvxjT)#roOqPD+FAuVUJgmNbheXFrF@!S4nvH!h+TLfCETl&t? zL(>L#*q@M4L*cmH-Jc5bF-s-@-BVIj;b8ZjR#jn_p{t%2O-~_oG$BVFIYFvFLf{Vr z8=6{-LSD@-W$fJ*Q!kCPXA&2Cs){M;JiCPq!;pi8;6vfy4HmEWu&$te(3YTnt=;gK z_AY&`GP8y_w~}=l&6zhIGA5T4f2j%xWEogjOaqcHQ!sGB*lT17h>flEZn-J{BWPj1 z(N#?bbQGgvUa!j<|NYFMMN2-NMwxo~){?3rr>~sez;eKfz#1)wy_1jd@lyGXz>38X zf4{Cu#64qtm`@fnAHFI#cgpa8w$k|-L1kSzD_?e1LP@EC8q>U=cg`sIn#H<+;2xH+ zbdA^5ZZa?D-X^ms`CuIG;nM$NZzU!7wzhrvy0KXO;xSpBSqA8vCg{~GW?++KiRIWl zia2>kqPlyX>jStSgWTkjF*6nT{@0DC7D@V{C(hc z851{T^m`dT^4PxlsN3vu8p^0m#AL03jmGSp<+uX5_R2t+pnvz|oKZ&S^J8MIbz0U7 z1-j}YDXVQTWc-daL6TV0=Ftm(+4h^bgc|?62C7EP3onq`@RD91#g8LTNH$AK8)-}y z_9-L|@8oeLRX_bkgdaGSQDS%KRT|Y~Om_v*Yn__q5~tpRC2K?g`L=G~;`NNzY)Db`5c+t7!}_1zw3f6O^^ zl3-*q{@LX4W{#z8@Mn#ltSt}`;|pyY2{`wG(=w5!MHt=387^;L07 zP=P_+cmqO4e7neP)5!?zt+lPqmcN-=MbLvA#WiNf8^Uk3Rx~rq2U(bdSD2-pw*}NW z@{HBrCilHv?<&g4d14^&OMC%|)XeF~-!tY)QY#nK7Z5B3tDBa3QIVrwl$B9Amapx| za4TzLKt;F@S)`>L#;?&(97rmzh3zkxyu+5Td_AOqO(@;GVx_L_kl3~^p(v+dpl4>> zV%c&HE1Uct+th))gZLF}fIKSV<43{efVsZsrKTGD?NWg#y3D)&?1{xJQWKT7n)tM@$>zjg=mp=QoT9?h(A_bP?7S;eZ119Cg-CWy%P+uxywh)62*Kve~m zeF*~aM6XWuE-9h+Mk!`<{A2Te9k~oOr}v-gF$dV{)B~YTegA<6EPv^xse46VMy4m< z?q=V@0{8#yt)7}QkuX{aY}G`osOM%2_bltMx^W@0N4+;k(r<~Z#b&(Gos-JWEm0r{ zLCISU*ej3ARctlONpA?=Nj@$w(?%Q4&>%uftKew4<=FAlVj`07^+j&)i(Sg|`jGT1 z1OoK2mKu24U!D$+A|k@d>R+)PAIx(njqRtOQ;iYh#K$Xsv-8{9`5OE(h@op>i7D_w zT>X5>7_ty6O?NgBLc&%*UJ%k{ zOaTx*Z;7AKmnBth69W7SuERWQ7uCj4*O<(5^FGxFN4YHx#NQK zE4}vn^Mke14)L3xY;ZdNpy@0b<~+qIUu(yF7d~s!#p7F;!XEUEu7{GDd{!qWRgE)i z-AivtSruW7BskK4dBT8XfQuSXBz|Xe*|ikr-=XE+r44nP&ZPHV7wOn}fQ`6K4m5hS zKeY3{ZpYQxB3jVz;On#(Jxd$G*5wAHfKsg7>dR^6@syh@dT7QK>wOdNb$>#yxBu|8 z0G!2f2GeiiG4=ggZT;w%HB)PAFX)7+(a((nw#HoyKuIT^A`{FB#T6LLlR4J5mgLT4 zKjJX|(O2|anU1cF%+g!43+8p=n?17JjO#V9;`e@BC&^=@CNx_WyY}Jms9MAj_+IrSPL88ueO~0W^zH6PP%BH@e9S^k0xjCJ7S@y65W#%jjVD#84q@uF zd9Mn)KSK#uE->8cYDKQOZ*(2=UcNQGv0F&=n>!Wczg<@GVz2In#i3hc_lp;GG3YV( z9{4WPrR1+o^G8fLd(y(3dL5G52*|(}uuA9!IT`1fbpfn@Q}7-0H9W=7AJ@UiVb8(` z3wa=JtY;*8b4F1R;PB6 zdPyXUBnZTu3cG(tt4dA7G?05A6&B?$NiZsVd4#O;C*O9^f?6bfch(OIQ1(i?p%1Pv zj(s2veP$o6SwAfm$C1^gFM7Bjog$;9G4>Fb=v?BfcH4QfB+kVBew@iG-@9*U#gv~K zZ^!3I7+2Jk`V1;U$Z-$<@^&+Pc38o*Ipj~0rPy_%{r#8f9b^yII#8m(txlVsU{uMe z-%9^J1NWhztJ(jWD(k&R8c*U&_csUx++kZu&=nvtC6_;X(U7v7N1Bh_u6d*jo_{S; z?>;4R;PQ-e%zNpat<>ZG^NTt~32uu;I=cBrYuvtR*DCnVF%62%)g9_PP7d;h*)Fcv zahSN1-ek_QswwJ%bA->YVzu=3HG5*gN1{l@0Hy}i!bfFD1(_4tC~M^-Em&kXOOeyIeW`8qIsSlQsG3Wsn46K_k_;A}bh|Q8UdT_jpG& zF6D4`#hA`(J1V^p>o|P)X1jeZunCb>&*VH6ykOedL7tX5TK)Tl9|$*@7`n-~q%CZv z81H1ax~$uDf?;0z#)KaBvg`7wLtMeAlMfwS`O-=S&b!%oC+M>*!+y`T<@2X4@BS-^ zLoZjeL8K}7tONz~iDn1ui-X89OCcT~V)nYq9}fm8i(pru`*Ja67*Du75!BfbsmJ|_ zmjTQJ+*(j1_%AxT`b&M*{R=SuD+vHOrl$tqXe^@qZjx62mmyKN>WTrQk`uCt1SIS|_~ToUxy zTe07OXwz2Ai~=B_t?PGalYo<69xzREF2soXeF2lzSjqhgy~EZFK$5Onu6a6maY*_I zpg`q+{t!(uIb>7_x#L$^>0NK)Y*I&FFunV_mxseX-^hK;VP(nOd>?8E>q%y5j1V z3^V20u%6t^F`)_SfG+IH;aL6J`XgMTEWf3Eo3~}D)Yk3v+p63aFiH845~o!0{)sqC zaQ-f_I$N-QYXqhK`T-|aEML5T&nz#_*l8MQjU3`rX|7?+g>-*xW;Pr2+wB4&^ZxAtix&5Am{yvDd%^(7ZZ=JtFm0V?yh6I{dn?k(Z#{f z{fi=73yvg{)%_x0NjVXNZsB7IlR-dJDO0QMzzaQ>T(0$r+a|Rcq;pv!FWOlR2_QY&9EpT@^)U##PSOu@J zZ&~rZhgtBnsqn^l) zH4cDNqB?`BUtaJ7D0)2sC%y*jnd?ljd0LWN-*^GC8^7d(GOt))Lip}~Qa;>G!6)y1 zh57RAtqYJahwC4hH?&vvlqq+bC-)CGS7bgctc^`?5w??1XYY2!o|QlQ=N0mLoQQ7m){eI^RmEBj+3bU%@nD0>R9 z>s<8qa`SNQ61T~MxHlw?Zu>~5`XyI331PJs_%N^DHq=bfGxE!zPlLT(uvB8od}sqyjGwo@sts4kS4 z6GytX7_a-4>Ze`H2Db=q^`$(eQy-@-c+Nn#^U~8A9=a|_!PlXF(D6#*wx@v3xIdw5 znH8S&c#dPN^*NGMpphs5)e(g zZlz(7Odyc2v-AsdeEr~7eQDf79}5WJ=~fhen$~wNbmJxO70E8TKevt3m{#ew_lDb3 z;FyEf-D!s+5BtZ$32Y)(N=)ol$mFOpR#fZKJ#!t)`m0eb*z=)*y0G{Gzs-dvp&*iN zD@0pY`%>*xCE|xr%*>vd3deDb`kD6xBX5w8qV*fn2Q=)U?rca!`Oa)tv2>NJnEzK* ziY!DyL*b!K;>gl{{{7IEF>gchTJ4?+WMNpN^_e(QSCq4bBFjkY!xPSf-!j7J*zrxr zBVbOd`{@+|C)5)Y*(H}Vz8^ePfnN6g-&k1Svc7lr%#b6twTU&RF;fM4Urf2-zXQD5> zY~WGnU||-w{)a46#7fyxY_}aMb!6hdglV9bMnjfkjBEyK|pM z^LFIhIRKx;*JawIvvYtAWp*%s|HDB9I~_;!^0~a3#OzY^apiu&!>FsB8?y<6r`R

TQC|751koWzKj1H z;-?L+vKx*S`n%kJ&7{QPBd?5S_HGX$Q7nVj%_=wl6o#Gcmeei86hdO@HPLm$lsEh@ zEge^n*_#MC&WZ|^-JMxAjrvV6Dk(+9++O&e;IAVC+5(wD zLGme&Sn1(*JqpXf`$IH=^r77U>O{=`f(N)vp8;C7)wi1>d$%V*R?xpDt#sj=6XW)f zbC>7jvn@-)lCzqWLeD$d55~|(gASUmqp8bL`wX_s9^GFa?+@H-sGtR27P94$R!R+u z%&;WV$m&Yv{f0n0z=|XfVcjx09Q5Q;AP1R_m8Tr1F=_ZbEX-_l?(~vIhCn6}ewgAaR6R1lnQn8wbzS%nd$cWizW>Ce zx21(<7Zl&^d}SdhchK@O)Fefp-$oo$FiJwtJ;(57VDV3Hu=d|T6RP8WaNjy!G5M>&2@D~w(s#5>Go96DjYYhvzHnj9E zCXeGv^&FysPIcBpUoYdZ_(&D8{C2J}y%c3T&+Cz+2b%$Xwl#gYD*70$+&l{u59J>V zuL1X!V1y`oK3LT{J<3~UGxKSgzZ`L3n<+`QEMNC&irZILoP_PD23$R3*j zt5?btJv8&3r!96a{y_ip;p+?5uf{l{tlZq!}v z@~+%8n$dvykB_zUv86R`kF^gija^f7{WBA@b02YW!9L|KEXXS~nlR0;6?CfF4K8&V zDQhY_izauHR0vLJ6+>msngYqew>UGfpK}_}c8YYuKvQ~WHXhu&Js0~L-qLT4%LZ>| zyvQ))>2&Ry<3SBJn)mhL6aF;4K${E zrH11^PqZ?7{hJ-&NCjf1dGOPy?|(>k-VkjY-Cd*`V}{sH9LsFe_`d(yI6~NKa!yh` zV$;{Z#iWk=x&?vfJbKq##)u)WsEA?Ksd_pQf7)X6mb26lOnbJgB4VwpnxXLaW@qls zMF00E*bHqa`*&!$@eXhP7^TMcb6l6`aZxS?C4)0BA9mFfrw(rn6MdM*pW({!HouKneG21!Eh(~@okSk-YZTzv2r@<9w_Ci!@$ z;7)|xlG7U_Qe5N`H_H1fxJvK&$0-k)WPbM>F#fF_&yRRMekYAtWgBr-c=eNZo1k!J zXyNIvTE@oO>xRpSedVO12n;kd4P@ixQzjGh-rT>NkfgQFRfm^@m>sC&pw$o6_W8H4 z{*-B4y$jx#d}4O(Rh0L3WgV&zmLVGVR=<$fX+!7Tqxaa)XE&SYu97ih%U9m=ZAqE-D z`WSEAo0eI+<|S7hq9}X7;NGn1#40*6~5X5EG5_yd#J|Tg@vFnVp`Gt;Z zM>Q5+A_&GU;$l&^T&Z`!o^=+!j%XvOy^@5k5St#b($+e4D zu{&+&@1|gu_C);ZzEx|ecCxX3etW~}>=8WOctOpu_=W(dkmZyRO0481xA^jFM{?25 zL0p!Wj+V>LLPrcMi%IR|n5QA!q;~rmICQDO(EaT?S&QE}0ll3xdYJx2vcIUZ()1E^jouy^r1&tyzDSa;%!{QT$p(3MRMjVT(g7lO(squ z&##Ri3Wh2x-S%`O=ilMCZ6#nN#^rLpUvHEoD5SqTE~+XZUX*_t3nfwHIE;+7B~!J~ zS;^kSIUnpZ`Q82`7grhdfmM#v5%1-1MN@gNA$V!~U58&Slk)J>ck$XmVH{B`W&Xa8 zzf9~L#4sK!x`np4v+iQl$+x+YZ{v~8lWZ;ZQG#_K_y!O07-2iEae9z3$wRJ2A80W2 zg=}EGw%z83OglDC(*3@7tjiN(uCp59n(TRJTwf8HsF{aej^@I-cVY=uiW(pP7<}Pu zheVhY9~T#7rg2^wuW_}}WX6_5vxP6lKRWp#Q!STYT<)>rzO%ER^PMW=P>l5O`5bo( zze98>i53N6On}`! z?Np&fOp(QAJC8iWjl6P)mOg$^f@UfXmUcs!EE9i=lZ&M@KLnlu1+Yh{ zn2U%J|HA0z3oxRHPpUuv9gNtaySN=Ns?3A0M~&2l&y^8`(FZMinaW;L&57Bu-xyQ$ z?#Uz@kK%WaPV2|yA0|uRe2y1I3{4u;un#^;;p04`l7K{I+-6P08y`oM3ORlL*KTw}MmD2%# z!T8Eibd;}^G}kN=Q%&bX4s1pd$?gcz#EIie=}k>pRY_$c%5Cn&8FPE^QK;qUTtSgL zWz25-?#fyMsfcK~!nm}uGWM$y6L^Qw;CS{FXaWwiv+LvnQ|Qss_j*36$$XAmM<#lY zNFHK6Wbmfhua0p*O7ylKiW&nAYUn=L`PVG%bqqw^d+6SuELXR~KxV>}9&DxE!>8Tn zmrbwTJUU(hg;eUd;${-_P6sule4kN|B80~h&jOEaKXVww#;g}+#5mcf_S2ktuw6NY`@*cdR+#5QV=NGNQ@;+ zl5q~U(lS~=ccM;cTyj4n5UdVe(Zfj=rMsG&Q@V^xXLf2+E+gBS9;~dO)i#$sM4dHx z(~ne6K7z_vP(bPH`C+1YsnZm-mH?`CjSML76BQ0oW$rXX+5U$qwifq(vx;d}*WRLf-?ULTos z1CM~kD8$armfv^6-PV1}7Ga?qEya68hYRBmKCGW=2&z4hMnTf|QcoxeFCN8sT`G5`hw{AwN;{85ufQzUpli$bI`*bQFejD;=s`dGOq;~usEK`5;$o&7E z`p15L`?qN3ALRb_leRZN&ftHB*@GXPJ^ERI=>IR|3?9;C-V4(CzL*q3-9Ah~?SA3i H$1na5CsAV_ literal 0 HcmV?d00001 diff --git a/third_party/quic-go/buffer_pool.go b/third_party/quic-go/buffer_pool.go new file mode 100644 index 0000000..34f3d1c --- /dev/null +++ b/third_party/quic-go/buffer_pool.go @@ -0,0 +1,92 @@ +package quic + +import ( + "sync" + + "github.com/apernet/quic-go/internal/protocol" +) + +type packetBuffer struct { + Data []byte + + // refCount counts how many packets Data is used in. + // It doesn't support concurrent use. + // It is > 1 when used for coalesced packet. + refCount int +} + +// Split increases the refCount. +// It must be called when a packet buffer is used for more than one packet, +// e.g. when splitting coalesced packets. +func (b *packetBuffer) Split() { + b.refCount++ +} + +// Decrement decrements the reference counter. +// It doesn't put the buffer back into the pool. +func (b *packetBuffer) Decrement() { + b.refCount-- + if b.refCount < 0 { + panic("negative packetBuffer refCount") + } +} + +// MaybeRelease puts the packet buffer back into the pool, +// if the reference counter already reached 0. +func (b *packetBuffer) MaybeRelease() { + // only put the packetBuffer back if it's not used any more + if b.refCount == 0 { + b.putBack() + } +} + +// Release puts back the packet buffer into the pool. +// It should be called when processing is definitely finished. +func (b *packetBuffer) Release() { + b.Decrement() + if b.refCount != 0 { + panic("packetBuffer refCount not zero") + } + b.putBack() +} + +// Len returns the length of Data +func (b *packetBuffer) Len() protocol.ByteCount { return protocol.ByteCount(len(b.Data)) } +func (b *packetBuffer) Cap() protocol.ByteCount { return protocol.ByteCount(cap(b.Data)) } + +func (b *packetBuffer) putBack() { + if cap(b.Data) == protocol.MaxPacketBufferSize { + bufferPool.Put(b) + return + } + if cap(b.Data) == protocol.MaxLargePacketBufferSize { + largeBufferPool.Put(b) + return + } + panic("putPacketBuffer called with packet of wrong size!") +} + +var bufferPool, largeBufferPool sync.Pool + +func getPacketBuffer() *packetBuffer { + buf := bufferPool.Get().(*packetBuffer) + buf.refCount = 1 + buf.Data = buf.Data[:0] + return buf +} + +func getLargePacketBuffer() *packetBuffer { + buf := largeBufferPool.Get().(*packetBuffer) + buf.refCount = 1 + buf.Data = buf.Data[:0] + return buf +} + +func init() { + bufferPool.New = func() any { + return &packetBuffer{Data: make([]byte, 0, protocol.MaxPacketBufferSize)} + } + largeBufferPool.New = func() any { + return &packetBuffer{Data: make([]byte, 0, protocol.MaxLargePacketBufferSize)} + } +} diff --git a/third_party/quic-go/buffer_pool_test.go b/third_party/quic-go/buffer_pool_test.go new file mode 100644 index 0000000..4f33c72 --- /dev/null +++ b/third_party/quic-go/buffer_pool_test.go @@ -0,0 +1,44 @@ +package quic + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestBufferPoolSizes(t *testing.T) { + buf1 := getPacketBuffer() + require.Equal(t, protocol.MaxPacketBufferSize, cap(buf1.Data)) + require.Zero(t, buf1.Len()) + buf1.Data = append(buf1.Data, []byte("foobar")...) + require.Equal(t, protocol.ByteCount(6), buf1.Len()) + + buf2 := getLargePacketBuffer() + require.Equal(t, protocol.MaxLargePacketBufferSize, cap(buf2.Data)) + require.Zero(t, buf2.Len()) +} + +func TestBufferPoolRelease(t *testing.T) { + buf1 := getPacketBuffer() + buf1.Release() + // panics if released twice + require.Panics(t, func() { buf1.Release() }) + + // panics if wrong-sized buffers are passed + buf2 := getLargePacketBuffer() + buf2.Data = make([]byte, 10) // replace the underlying slice + require.Panics(t, func() { buf2.Release() }) +} + +func TestBufferPoolSplitting(t *testing.T) { + buf := getPacketBuffer() + buf.Split() + buf.Split() + // now we have 3 parts + buf.Decrement() + buf.Decrement() + buf.Decrement() + require.Panics(t, func() { buf.Decrement() }) +} diff --git a/third_party/quic-go/chrome_parrot.go b/third_party/quic-go/chrome_parrot.go new file mode 100644 index 0000000..4bcbeb4 --- /dev/null +++ b/third_party/quic-go/chrome_parrot.go @@ -0,0 +1,59 @@ +package quic + +import ( + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +// Values a Chrome-parroting client pins. Advertising Chrome's transport +// parameter shape while carrying quic-go's own numbers would defeat the point. +const ( + chromeMaxIdleTimeout = 30 * time.Second + chromeInitialMaxStreamData = 6291456 + chromeInitialMaxData = 15728640 + chromeMaxIncomingStreams = 100 + chromeMaxIncomingUniStreams = 103 + // UDP payload size for the packets we send. + chromeInitialPacketSize = 1250 + // max_udp_payload_size: what we are willing to receive, which differs from + // what we send. + chromeMaxUDPPayloadSize = 1472 + // DATAGRAM support is always advertised, so ChromeParrot forces datagrams on + // rather than leave the fingerprint a parameter short. + chromeMaxDatagramFrameSize = 65536 +) + +// chromeParrotTransportParameters adjusts the transport parameters advertised by +// a Chrome-parroting client. The omitted ones are reset to their protocol +// defaults, so what we imply by omission matches how we actually behave. +func chromeParrotTransportParameters(params *wire.TransportParameters) { + params.MaxUDPPayloadSize = chromeMaxUDPPayloadSize + params.MaxDatagramFrameSize = chromeMaxDatagramFrameSize + params.MaxAckDelay = protocol.DefaultMaxAckDelay + params.AckDelayExponent = protocol.DefaultAckDelayExponent + params.DisableActiveMigration = false + params.ActiveConnectionIDLimit = protocol.DefaultActiveConnectionIDLimit + params.ChromeFingerprint = true +} + +// ZeroLengthConnectionIDGenerator generates zero-length connection IDs, as a +// Chrome-parroting client needs. Set it on a Transport: +// +// tr := &quic.Transport{Conn: conn, ConnectionIDGenerator: quic.ZeroLengthConnectionIDGenerator{}} +// +// It must be chosen at the Transport level, not per-dial: the Transport parses +// every incoming packet's destination connection ID at one fixed length, so +// overriding it for a single connection breaks routing. +// +// Handlers are keyed by source connection ID, so a Transport using this carries +// one connection at a time. That suits a client with one connection per socket, +// not a server or a multiplexing dialer. +type ZeroLengthConnectionIDGenerator struct{} + +func (ZeroLengthConnectionIDGenerator) GenerateConnectionID() (ConnectionID, error) { + return ConnectionID{}, nil +} + +func (ZeroLengthConnectionIDGenerator) ConnectionIDLen() int { return 0 } diff --git a/third_party/quic-go/client.go b/third_party/quic-go/client.go new file mode 100644 index 0000000..a69a390 --- /dev/null +++ b/third_party/quic-go/client.go @@ -0,0 +1,109 @@ +package quic + +import ( + "context" + "crypto/tls" + "errors" + "net" + + "github.com/apernet/quic-go/internal/protocol" +) + +// make it possible to mock connection ID for initial generation in the tests +var generateConnectionIDForInitial = protocol.GenerateConnectionIDForInitial + +// DialAddr establishes a new QUIC connection to a server. +// It resolves the address, and then creates a new UDP connection to dial the QUIC server. +// When the QUIC connection is closed, this UDP connection is closed. +// See [Dial] for more details. +func DialAddr(ctx context.Context, addr string, tlsConf *tls.Config, conf *Config) (*Conn, error) { + udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4zero, Port: 0}) + if err != nil { + return nil, err + } + udpAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return nil, err + } + tr, err := setupTransport(udpConn, tlsConf, true) + if err != nil { + return nil, err + } + conn, err := tr.dial(ctx, udpAddr, addr, tlsConf, conf, false) + if err != nil { + tr.Close() + return nil, err + } + return conn, nil +} + +// DialAddrEarly establishes a new 0-RTT QUIC connection to a server. +// See [DialAddr] for more details. +func DialAddrEarly(ctx context.Context, addr string, tlsConf *tls.Config, conf *Config) (*Conn, error) { + udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4zero, Port: 0}) + if err != nil { + return nil, err + } + udpAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return nil, err + } + tr, err := setupTransport(udpConn, tlsConf, true) + if err != nil { + return nil, err + } + conn, err := tr.dial(ctx, udpAddr, addr, tlsConf, conf, true) + if err != nil { + tr.Close() + return nil, err + } + return conn, nil +} + +// DialEarly establishes a new 0-RTT QUIC connection to a server using a net.PacketConn. +// See [Dial] for more details. +func DialEarly(ctx context.Context, c net.PacketConn, addr net.Addr, tlsConf *tls.Config, conf *Config) (*Conn, error) { + dl, err := setupTransport(c, tlsConf, false) + if err != nil { + return nil, err + } + conn, err := dl.DialEarly(ctx, addr, tlsConf, conf) + if err != nil { + dl.Close() + return nil, err + } + return conn, nil +} + +// Dial establishes a new QUIC connection to a server using a net.PacketConn. +// If the PacketConn satisfies the [OOBCapablePacketConn] interface (as a [net.UDPConn] does), +// ECN and packet info support will be enabled. In this case, ReadMsgUDP and WriteMsgUDP +// will be used instead of ReadFrom and WriteTo to read/write packets. +// The [tls.Config] must define an application protocol (using tls.Config.NextProtos). +// +// This is a convenience function. More advanced use cases should instantiate a [Transport], +// which offers configuration options for a more fine-grained control of the connection establishment, +// including reusing the underlying UDP socket for multiple QUIC connections. +func Dial(ctx context.Context, c net.PacketConn, addr net.Addr, tlsConf *tls.Config, conf *Config) (*Conn, error) { + dl, err := setupTransport(c, tlsConf, false) + if err != nil { + return nil, err + } + conn, err := dl.Dial(ctx, addr, tlsConf, conf) + if err != nil { + dl.Close() + return nil, err + } + return conn, nil +} + +func setupTransport(c net.PacketConn, tlsConf *tls.Config, createdPacketConn bool) (*Transport, error) { + if tlsConf == nil { + return nil, errors.New("quic: tls.Config not set") + } + return &Transport{ + Conn: c, + createdConn: createdPacketConn, + isSingleUse: true, + }, nil +} diff --git a/third_party/quic-go/client_test.go b/third_party/quic-go/client_test.go new file mode 100644 index 0000000..6dad6ad --- /dev/null +++ b/third_party/quic-go/client_test.go @@ -0,0 +1,105 @@ +package quic + +import ( + "context" + "crypto/tls" + "net" + "runtime" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestDial(t *testing.T) { + t.Run("Dial", func(t *testing.T) { + testDial(t, + func(ctx context.Context, addr net.Addr) error { + conn := newUDPConnLocalhost(t) + _, err := Dial(ctx, conn, addr, &tls.Config{InsecureSkipVerify: true}, nil) + return err + }, + false, + ) + }) + + t.Run("DialEarly", func(t *testing.T) { + testDial(t, + func(ctx context.Context, addr net.Addr) error { + conn := newUDPConnLocalhost(t) + _, err := DialEarly(ctx, conn, addr, &tls.Config{InsecureSkipVerify: true}, nil) + return err + }, + false, + ) + }) + + t.Run("DialAddr", func(t *testing.T) { + testDial(t, + func(ctx context.Context, addr net.Addr) error { + _, err := DialAddr(ctx, addr.String(), &tls.Config{InsecureSkipVerify: true}, nil) + return err + }, + true, + ) + }) + + t.Run("DialAddrEarly", func(t *testing.T) { + testDial(t, + func(ctx context.Context, addr net.Addr) error { + _, err := DialAddrEarly(ctx, addr.String(), &tls.Config{InsecureSkipVerify: true}, nil) + return err + }, + true, + ) + }) +} + +func testDial(t *testing.T, + dialFn func(context.Context, net.Addr) error, + shouldCloseConn bool, +) { + server := newUDPConnLocalhost(t) + + ctx, cancel := context.WithCancel(context.Background()) + errChan := make(chan error, 1) + go func() { errChan <- dialFn(ctx, server.LocalAddr()) }() + + server.SetReadDeadline(time.Now().Add(time.Second)) + _, addr, err := server.ReadFrom(make([]byte, 1500)) + require.NoError(t, err) + cancel() + select { + case err := <-errChan: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + if shouldCloseConn { + // The socket that the client used for dialing should be closed now. + // Binding to the same address would error if the address was still in use. + require.Eventually(t, func() bool { + conn, err := net.ListenUDP("udp", addr.(*net.UDPAddr)) + if err != nil { + return false + } + conn.Close() + return true + }, scaleDuration(200*time.Millisecond), scaleDuration(10*time.Millisecond)) + require.False(t, areTransportsRunning()) + return + } + + // The socket that the client used for dialing should not be closed now. + // Binding to the same address will error if the address was still in use. + _, err = net.ListenUDP("udp", addr.(*net.UDPAddr)) + require.Error(t, err) + if runtime.GOOS == "windows" { + require.ErrorContains(t, err, "bind: Only one usage of each socket address") + } else { + require.ErrorContains(t, err, "address already in use") + } + + require.False(t, areTransportsRunning()) +} diff --git a/third_party/quic-go/closed_conn.go b/third_party/quic-go/closed_conn.go new file mode 100644 index 0000000..6f0b15c --- /dev/null +++ b/third_party/quic-go/closed_conn.go @@ -0,0 +1,58 @@ +package quic + +import ( + "math/bits" + "net" + "sync/atomic" + + "github.com/apernet/quic-go/internal/utils" +) + +// A closedLocalConn is a connection that we closed locally. +// When receiving packets for such a connection, we need to retransmit the packet containing the CONNECTION_CLOSE frame, +// with an exponential backoff. +type closedLocalConn struct { + counter atomic.Uint32 + logger utils.Logger + + sendPacket func(net.Addr, packetInfo) +} + +var _ packetHandler = &closedLocalConn{} + +// newClosedLocalConn creates a new closedLocalConn and runs it. +func newClosedLocalConn(sendPacket func(net.Addr, packetInfo), logger utils.Logger) packetHandler { + return &closedLocalConn{ + sendPacket: sendPacket, + logger: logger, + } +} + +func (c *closedLocalConn) handlePacket(p receivedPacket) { + n := c.counter.Add(1) + // exponential backoff + // only send a CONNECTION_CLOSE for the 1st, 2nd, 4th, 8th, 16th, ... packet arriving + if bits.OnesCount32(n) != 1 { + return + } + c.logger.Debugf("Received %d packets after sending CONNECTION_CLOSE. Retransmitting.", n) + c.sendPacket(p.remoteAddr, p.info) +} + +func (c *closedLocalConn) destroy(error) {} +func (c *closedLocalConn) closeWithTransportError(TransportErrorCode) {} + +// A closedRemoteConn is a connection that was closed remotely. +// For such a connection, we might receive reordered packets that were sent before the CONNECTION_CLOSE. +// We can just ignore those packets. +type closedRemoteConn struct{} + +var _ packetHandler = &closedRemoteConn{} + +func newClosedRemoteConn() packetHandler { + return &closedRemoteConn{} +} + +func (c *closedRemoteConn) handlePacket(receivedPacket) {} +func (c *closedRemoteConn) destroy(error) {} +func (c *closedRemoteConn) closeWithTransportError(TransportErrorCode) {} diff --git a/third_party/quic-go/closed_conn_test.go b/third_party/quic-go/closed_conn_test.go new file mode 100644 index 0000000..fae8c13 --- /dev/null +++ b/third_party/quic-go/closed_conn_test.go @@ -0,0 +1,34 @@ +package quic + +import ( + "net" + "testing" + + "github.com/apernet/quic-go/internal/utils" + + "github.com/stretchr/testify/require" +) + +func TestClosedLocalConnection(t *testing.T) { + written := make(chan net.Addr, 1) + conn := newClosedLocalConn(func(addr net.Addr, _ packetInfo) { written <- addr }, utils.DefaultLogger) + addr := &net.UDPAddr{IP: net.IPv4(127, 1, 2, 3), Port: 1337} + for i := 1; i <= 20; i++ { + conn.handlePacket(receivedPacket{remoteAddr: addr}) + if i == 1 || i == 2 || i == 4 || i == 8 || i == 16 { + select { + case gotAddr := <-written: + require.Equal(t, addr, gotAddr) // receive the CONNECTION_CLOSE + default: + t.Fatal("expected to receive address") + } + } else { + select { + case gotAddr := <-written: + t.Fatalf("unexpected address received: %v", gotAddr) + default: + // Nothing received, which is expected + } + } + } +} diff --git a/third_party/quic-go/codecov.yml b/third_party/quic-go/codecov.yml new file mode 100644 index 0000000..55d8a35 --- /dev/null +++ b/third_party/quic-go/codecov.yml @@ -0,0 +1,23 @@ +coverage: + round: nearest + ignore: + - http3/gzip_reader.go + - example/ + - interop/ + - internal/handshake/cipher_suite.go + - internal/mocks/ + - internal/utils/linkedlist/linkedlist.go + - internal/testdata + - testutils/ + - fuzzing/ + - metrics/ + status: + project: + default: + threshold: 0.5 + patch: false +flags: + clusterfuzz-lite-batch: + joined: false + clusterfuzz: + joined: false diff --git a/third_party/quic-go/config.go b/third_party/quic-go/config.go new file mode 100644 index 0000000..daf8703 --- /dev/null +++ b/third_party/quic-go/config.go @@ -0,0 +1,153 @@ +package quic + +import ( + "fmt" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// Clone clones a Config. +func (c *Config) Clone() *Config { + copy := *c + return © +} + +func (c *Config) handshakeTimeout() time.Duration { + return 2 * c.HandshakeIdleTimeout +} + +func (c *Config) maxRetryTokenAge() time.Duration { + return c.handshakeTimeout() +} + +func validateConfig(config *Config) error { + if config == nil { + return nil + } + const maxStreams = 1 << 60 + if config.MaxIncomingStreams > maxStreams { + config.MaxIncomingStreams = maxStreams + } + if config.MaxIncomingUniStreams > maxStreams { + config.MaxIncomingUniStreams = maxStreams + } + if config.MaxStreamReceiveWindow > quicvarint.Max { + config.MaxStreamReceiveWindow = quicvarint.Max + } + if config.MaxConnectionReceiveWindow > quicvarint.Max { + config.MaxConnectionReceiveWindow = quicvarint.Max + } + if config.InitialPacketSize > 0 && config.InitialPacketSize < protocol.MinInitialPacketSize { + config.InitialPacketSize = protocol.MinInitialPacketSize + } + if config.InitialPacketSize > protocol.MaxPacketBufferSize { + config.InitialPacketSize = protocol.MaxPacketBufferSize + } + // check that all QUIC versions are actually supported + for _, v := range config.Versions { + if !protocol.IsValidVersion(v) { + return fmt.Errorf("invalid QUIC version: %s", v) + } + } + return nil +} + +// populateConfig populates fields in the quic.Config with their default values, if none are set +// it may be called with nil +func populateConfig(config *Config) *Config { + if config == nil { + config = &Config{} + } + versions := config.Versions + if len(versions) == 0 { + versions = protocol.SupportedVersions + } + handshakeIdleTimeout := protocol.DefaultHandshakeIdleTimeout + if config.HandshakeIdleTimeout != 0 { + handshakeIdleTimeout = config.HandshakeIdleTimeout + } + idleTimeout := protocol.DefaultIdleTimeout + if config.MaxIdleTimeout != 0 { + idleTimeout = config.MaxIdleTimeout + } + initialStreamReceiveWindow := config.InitialStreamReceiveWindow + if initialStreamReceiveWindow == 0 { + initialStreamReceiveWindow = protocol.DefaultInitialMaxStreamData + } + maxStreamReceiveWindow := config.MaxStreamReceiveWindow + if maxStreamReceiveWindow == 0 { + maxStreamReceiveWindow = protocol.DefaultMaxReceiveStreamFlowControlWindow + } + initialConnectionReceiveWindow := config.InitialConnectionReceiveWindow + if initialConnectionReceiveWindow == 0 { + initialConnectionReceiveWindow = protocol.DefaultInitialMaxData + } + maxConnectionReceiveWindow := config.MaxConnectionReceiveWindow + if maxConnectionReceiveWindow == 0 { + maxConnectionReceiveWindow = protocol.DefaultMaxReceiveConnectionFlowControlWindow + } + maxIncomingStreams := config.MaxIncomingStreams + if maxIncomingStreams == 0 { + maxIncomingStreams = protocol.DefaultMaxIncomingStreams + } else if maxIncomingStreams < 0 { + maxIncomingStreams = 0 + } + maxIncomingUniStreams := config.MaxIncomingUniStreams + if maxIncomingUniStreams == 0 { + maxIncomingUniStreams = protocol.DefaultMaxIncomingUniStreams + } else if maxIncomingUniStreams < 0 { + maxIncomingUniStreams = 0 + } + initialPacketSize := config.InitialPacketSize + if initialPacketSize == 0 { + initialPacketSize = protocol.InitialPacketSize + } + enableDatagrams := config.EnableDatagrams + omitMaxDatagramFrameSize := config.OmitMaxDatagramFrameSize + if config.ChromeParrot { + // Chrome always advertises DATAGRAM support, so enable it and never omit + // the transport parameter; leaving it out would be one parameter short of + // Chrome's set. + enableDatagrams = true + omitMaxDatagramFrameSize = false + // Chrome pins these, so anything the caller asked for is overridden. + idleTimeout = chromeMaxIdleTimeout + initialStreamReceiveWindow = chromeInitialMaxStreamData + initialConnectionReceiveWindow = chromeInitialMaxData + maxIncomingStreams = chromeMaxIncomingStreams + maxIncomingUniStreams = chromeMaxIncomingUniStreams + initialPacketSize = chromeInitialPacketSize + // The auto-tuning ceilings must not sit below the starting windows. + maxStreamReceiveWindow = max(maxStreamReceiveWindow, initialStreamReceiveWindow) + maxConnectionReceiveWindow = max(maxConnectionReceiveWindow, initialConnectionReceiveWindow) + } + + return &Config{ + GetConfigForClient: config.GetConfigForClient, + Versions: versions, + HandshakeIdleTimeout: handshakeIdleTimeout, + MaxIdleTimeout: idleTimeout, + KeepAlivePeriod: config.KeepAlivePeriod, + InitialStreamReceiveWindow: initialStreamReceiveWindow, + MaxStreamReceiveWindow: maxStreamReceiveWindow, + InitialConnectionReceiveWindow: initialConnectionReceiveWindow, + MaxConnectionReceiveWindow: maxConnectionReceiveWindow, + AllowConnectionWindowIncrease: config.AllowConnectionWindowIncrease, + MaxIncomingStreams: maxIncomingStreams, + MaxIncomingUniStreams: maxIncomingUniStreams, + TokenStore: config.TokenStore, + EnableDatagrams: enableDatagrams, + OmitMaxDatagramFrameSize: omitMaxDatagramFrameSize, + AssumePeerMaxDatagramFrameSize: config.AssumePeerMaxDatagramFrameSize, + InitialPacketSize: initialPacketSize, + DisablePathMTUDiscovery: config.DisablePathMTUDiscovery, + EnableStreamResetPartialDelivery: config.EnableStreamResetPartialDelivery, + Allow0RTT: config.Allow0RTT, + Tracer: config.Tracer, + MaxDatagramFrameSize: config.MaxDatagramFrameSize, + DisablePathManager: config.DisablePathManager, + ChromeParrot: config.ChromeParrot, + } +} diff --git a/third_party/quic-go/config_test.go b/third_party/quic-go/config_test.go new file mode 100644 index 0000000..d08aa2c --- /dev/null +++ b/third_party/quic-go/config_test.go @@ -0,0 +1,250 @@ +package quic + +import ( + "context" + "reflect" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestConfigValidation(t *testing.T) { + t.Run("nil config", func(t *testing.T) { + require.NoError(t, validateConfig(nil)) + }) + + t.Run("config with a few values set", func(t *testing.T) { + conf := populateConfig(&Config{ + MaxIncomingStreams: 5, + MaxStreamReceiveWindow: 10, + }) + require.NoError(t, validateConfig(conf)) + require.Equal(t, int64(5), conf.MaxIncomingStreams) + require.Equal(t, uint64(10), conf.MaxStreamReceiveWindow) + }) + + t.Run("stream limits", func(t *testing.T) { + conf := &Config{ + MaxIncomingStreams: 1<<60 + 1, + MaxIncomingUniStreams: 1<<60 + 2, + } + require.NoError(t, validateConfig(conf)) + require.Equal(t, int64(1<<60), conf.MaxIncomingStreams) + require.Equal(t, int64(1<<60), conf.MaxIncomingUniStreams) + }) + + t.Run("flow control windows", func(t *testing.T) { + conf := &Config{ + MaxStreamReceiveWindow: quicvarint.Max + 1, + MaxConnectionReceiveWindow: quicvarint.Max + 2, + } + require.NoError(t, validateConfig(conf)) + require.Equal(t, uint64(quicvarint.Max), conf.MaxStreamReceiveWindow) + require.Equal(t, uint64(quicvarint.Max), conf.MaxConnectionReceiveWindow) + }) + + t.Run("initial packet size", func(t *testing.T) { + // not set + conf := &Config{InitialPacketSize: 0} + require.NoError(t, validateConfig(conf)) + require.Zero(t, conf.InitialPacketSize) + + // too small + conf = &Config{InitialPacketSize: 10} + require.NoError(t, validateConfig(conf)) + require.Equal(t, uint16(1200), conf.InitialPacketSize) + + // too large + conf = &Config{InitialPacketSize: protocol.MaxPacketBufferSize + 1} + require.NoError(t, validateConfig(conf)) + require.Equal(t, uint16(protocol.MaxPacketBufferSize), conf.InitialPacketSize) + }) +} + +func TestConfigHandshakeIdleTimeout(t *testing.T) { + c := &Config{HandshakeIdleTimeout: time.Second * 11 / 2} + require.Equal(t, 11*time.Second, c.handshakeTimeout()) +} + +// chromeParrot is set separately because, unlike every other field here, it +// deliberately overrides other values in populateConfig. Callers that assert +// populateConfig is idempotent must leave it off. +func configWithNonZeroNonFunctionFields(t *testing.T, chromeParrot bool) *Config { + t.Helper() + c := &Config{} + v := reflect.ValueOf(c).Elem() + + typ := v.Type() + for i := 0; i < typ.NumField(); i++ { + f := v.Field(i) + if !f.CanSet() { + // unexported field; not cloned. + continue + } + + switch fn := typ.Field(i).Name; fn { + case "GetConfigForClient", "RequireAddressValidation", "GetLogWriter", "AllowConnectionWindowIncrease", "Tracer": + // Can't compare functions. + case "Versions": + f.Set(reflect.ValueOf([]Version{1, 2, 3})) + case "ConnectionIDLength": + f.Set(reflect.ValueOf(8)) + case "ConnectionIDGenerator": + f.Set(reflect.ValueOf(&protocol.DefaultConnectionIDGenerator{ConnLen: protocol.DefaultConnectionIDLength})) + case "HandshakeIdleTimeout": + f.Set(reflect.ValueOf(time.Second)) + case "MaxIdleTimeout": + f.Set(reflect.ValueOf(time.Hour)) + case "TokenStore": + f.Set(reflect.ValueOf(NewLRUTokenStore(2, 3))) + case "InitialStreamReceiveWindow": + f.Set(reflect.ValueOf(uint64(1234))) + case "MaxStreamReceiveWindow": + f.Set(reflect.ValueOf(uint64(9))) + case "InitialConnectionReceiveWindow": + f.Set(reflect.ValueOf(uint64(4321))) + case "MaxConnectionReceiveWindow": + f.Set(reflect.ValueOf(uint64(10))) + case "MaxIncomingStreams": + f.Set(reflect.ValueOf(int64(11))) + case "MaxIncomingUniStreams": + f.Set(reflect.ValueOf(int64(12))) + case "StatelessResetKey": + f.Set(reflect.ValueOf(&StatelessResetKey{1, 2, 3, 4})) + case "KeepAlivePeriod": + f.Set(reflect.ValueOf(time.Second)) + case "EnableDatagrams": + f.Set(reflect.ValueOf(true)) + case "OmitMaxDatagramFrameSize": + f.Set(reflect.ValueOf(true)) + case "AssumePeerMaxDatagramFrameSize": + f.Set(reflect.ValueOf(int64(1337))) + case "MaxDatagramFrameSize": + f.Set(reflect.ValueOf(int64(1200))) + case "DisablePathManager": + f.Set(reflect.ValueOf(true)) + case "DisableVersionNegotiationPackets": + f.Set(reflect.ValueOf(true)) + case "InitialPacketSize": + f.Set(reflect.ValueOf(uint16(1350))) + case "DisablePathMTUDiscovery": + f.Set(reflect.ValueOf(true)) + case "Allow0RTT": + f.Set(reflect.ValueOf(true)) + case "EnableStreamResetPartialDelivery": + f.Set(reflect.ValueOf(true)) + case "ChromeParrot": + f.Set(reflect.ValueOf(chromeParrot)) + default: + t.Fatalf("all fields must be accounted for, but saw unknown field %q", fn) + } + } + return c +} + +func TestConfigClone(t *testing.T) { + t.Run("function fields", func(t *testing.T) { + var calledAllowConnectionWindowIncrease, calledTracer bool + c1 := &Config{ + GetConfigForClient: func(info *ClientInfo) (*Config, error) { return nil, assert.AnError }, + AllowConnectionWindowIncrease: func(*Conn, uint64) bool { calledAllowConnectionWindowIncrease = true; return true }, + Tracer: func(context.Context, bool, ConnectionID) qlogwriter.Trace { + calledTracer = true + return nil + }, + } + c2 := c1.Clone() + c2.AllowConnectionWindowIncrease(nil, 1234) + require.True(t, calledAllowConnectionWindowIncrease) + _, err := c2.GetConfigForClient(&ClientInfo{}) + require.ErrorIs(t, err, assert.AnError) + c2.Tracer(context.Background(), true, protocol.ConnectionID{}) + require.True(t, calledTracer) + }) + + t.Run("non-function fields", func(t *testing.T) { + c := configWithNonZeroNonFunctionFields(t, true) + require.Equal(t, c, c.Clone()) + }) + + t.Run("returns a copy", func(t *testing.T) { + c1 := &Config{MaxIncomingStreams: 100} + c2 := c1.Clone() + c2.MaxIncomingStreams = 200 + require.EqualValues(t, 100, c1.MaxIncomingStreams) + }) +} + +func TestConfigDefaultValues(t *testing.T) { + // if set, the values should be copied + c := configWithNonZeroNonFunctionFields(t, false) + require.Equal(t, c, populateConfig(c)) + + // if not set, some fields use default values + c = populateConfig(&Config{}) + require.Equal(t, protocol.SupportedVersions, c.Versions) + require.Equal(t, protocol.DefaultHandshakeIdleTimeout, c.HandshakeIdleTimeout) + require.Equal(t, protocol.DefaultIdleTimeout, c.MaxIdleTimeout) + require.EqualValues(t, protocol.DefaultInitialMaxStreamData, c.InitialStreamReceiveWindow) + require.EqualValues(t, protocol.DefaultMaxReceiveStreamFlowControlWindow, c.MaxStreamReceiveWindow) + require.EqualValues(t, protocol.DefaultInitialMaxData, c.InitialConnectionReceiveWindow) + require.EqualValues(t, protocol.DefaultMaxReceiveConnectionFlowControlWindow, c.MaxConnectionReceiveWindow) + require.EqualValues(t, protocol.DefaultMaxIncomingStreams, c.MaxIncomingStreams) + require.EqualValues(t, protocol.DefaultMaxIncomingUniStreams, c.MaxIncomingUniStreams) + require.False(t, c.DisablePathMTUDiscovery) + require.Nil(t, c.GetConfigForClient) +} + +func TestConfigZeroLimits(t *testing.T) { + config := &Config{ + MaxIncomingStreams: -1, + MaxIncomingUniStreams: -1, + } + c := populateConfig(config) + require.Zero(t, c.MaxIncomingStreams) + require.Zero(t, c.MaxIncomingUniStreams) +} + +func TestConfigChromeParrotOverrides(t *testing.T) { + // ChromeParrot pins the values that show up in the transport parameters, so + // whatever the caller asked for must be discarded. + c := populateConfig(&Config{ + ChromeParrot: true, + MaxIdleTimeout: time.Hour, + InitialStreamReceiveWindow: 1234, + InitialConnectionReceiveWindow: 4321, + MaxIncomingStreams: 7, + MaxIncomingUniStreams: 9, + InitialPacketSize: 1350, + }) + + require.Equal(t, chromeMaxIdleTimeout, c.MaxIdleTimeout) + require.EqualValues(t, chromeInitialMaxStreamData, c.InitialStreamReceiveWindow) + require.EqualValues(t, chromeInitialMaxData, c.InitialConnectionReceiveWindow) + require.EqualValues(t, chromeMaxIncomingStreams, c.MaxIncomingStreams) + require.EqualValues(t, chromeMaxIncomingUniStreams, c.MaxIncomingUniStreams) + require.EqualValues(t, chromeInitialPacketSize, c.InitialPacketSize) + + // The auto-tuning ceilings must never end up below the starting windows. + require.GreaterOrEqual(t, c.MaxStreamReceiveWindow, c.InitialStreamReceiveWindow) + require.GreaterOrEqual(t, c.MaxConnectionReceiveWindow, c.InitialConnectionReceiveWindow) +} + +func TestConfigChromeParrotForcesDatagrams(t *testing.T) { + // Chrome always advertises max_datagram_frame_size, so ChromeParrot must turn + // datagrams on and clear Hysteria's OmitMaxDatagramFrameSize quirk; otherwise + // the transport parameter set comes out one short of Chrome's. + c := populateConfig(&Config{ + ChromeParrot: true, + EnableDatagrams: false, + OmitMaxDatagramFrameSize: true, + }) + require.True(t, c.EnableDatagrams) + require.False(t, c.OmitMaxDatagramFrameSize) +} diff --git a/third_party/quic-go/congestion/interface.go b/third_party/quic-go/congestion/interface.go new file mode 100644 index 0000000..1db771d --- /dev/null +++ b/third_party/quic-go/congestion/interface.go @@ -0,0 +1,67 @@ +package congestion + +import ( + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/monotime" +) + +type ( + ByteCount protocol.ByteCount + PacketNumber protocol.PacketNumber +) + +// Expose some constants from protocol that congestion control algorithms may need. +const ( + InitialPacketSize = protocol.InitialPacketSize + MinPacingDelay = protocol.MinPacingDelay + MaxPacketBufferSize = protocol.MaxPacketBufferSize + MinInitialPacketSize = protocol.MinInitialPacketSize + MaxCongestionWindowPackets = protocol.MaxCongestionWindowPackets + PacketsPerConnectionID = protocol.PacketsPerConnectionID +) + +type AckedPacketInfo struct { + PacketNumber PacketNumber + BytesAcked ByteCount + ReceivedTime monotime.Time +} + +type LostPacketInfo struct { + PacketNumber PacketNumber + BytesLost ByteCount +} + +type CongestionControl interface { + SetRTTStatsProvider(provider RTTStatsProvider) + TimeUntilSend(bytesInFlight ByteCount) monotime.Time + HasPacingBudget(now monotime.Time) bool + OnPacketSent(sentTime monotime.Time, bytesInFlight ByteCount, packetNumber PacketNumber, bytes ByteCount, isRetransmittable bool) + CanSend(bytesInFlight ByteCount) bool + MaybeExitSlowStart() + OnPacketAcked(number PacketNumber, ackedBytes ByteCount, priorInFlight ByteCount, eventTime monotime.Time) + OnCongestionEvent(number PacketNumber, lostBytes ByteCount, priorInFlight ByteCount) + OnRetransmissionTimeout(packetsRetransmitted bool) + SetMaxDatagramSize(size ByteCount) + InSlowStart() bool + InRecovery() bool + GetCongestionWindow() ByteCount +} + +type CongestionControlEx interface { + CongestionControl + OnCongestionEventEx(priorInFlight ByteCount, eventTime monotime.Time, ackedPackets []AckedPacketInfo, lostPackets []LostPacketInfo) +} + +type RTTStatsProvider interface { + MinRTT() time.Duration + LatestRTT() time.Duration + SmoothedRTT() time.Duration + MeanDeviation() time.Duration + MaxAckDelay() time.Duration + PTO(includeMaxAckDelay bool) time.Duration + UpdateRTT(sendDelta, ackDelay time.Duration) + SetMaxAckDelay(mad time.Duration) + SetInitialRTT(t time.Duration) +} diff --git a/third_party/quic-go/conn_id_generator.go b/third_party/quic-go/conn_id_generator.go new file mode 100644 index 0000000..47a32c2 --- /dev/null +++ b/third_party/quic-go/conn_id_generator.go @@ -0,0 +1,212 @@ +package quic + +import ( + "fmt" + "slices" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +type connRunnerCallbacks struct { + AddConnectionID func(protocol.ConnectionID) + RemoveConnectionID func(protocol.ConnectionID) + ReplaceWithClosed func([]protocol.ConnectionID, []byte, time.Duration) +} + +// The memory address of the Transport is used as the key. +type connRunners map[connRunner]connRunnerCallbacks + +func (cr connRunners) AddConnectionID(id protocol.ConnectionID) { + for _, c := range cr { + c.AddConnectionID(id) + } +} + +func (cr connRunners) RemoveConnectionID(id protocol.ConnectionID) { + for _, c := range cr { + c.RemoveConnectionID(id) + } +} + +func (cr connRunners) ReplaceWithClosed(ids []protocol.ConnectionID, b []byte, expiry time.Duration) { + for _, c := range cr { + c.ReplaceWithClosed(ids, b, expiry) + } +} + +type connIDToRetire struct { + t monotime.Time + connID protocol.ConnectionID +} + +type connIDGenerator struct { + generator ConnectionIDGenerator + highestSeq uint64 + connRunners connRunners + + activeSrcConnIDs map[uint64]protocol.ConnectionID + connIDsToRetire []connIDToRetire // sorted by t + initialClientDestConnID *protocol.ConnectionID // nil for the client + + statelessResetter *statelessResetter + + queueControlFrame func(wire.Frame) +} + +func newConnIDGenerator( + runner connRunner, + initialConnectionID protocol.ConnectionID, + initialClientDestConnID *protocol.ConnectionID, // nil for the client + statelessResetter *statelessResetter, + callbacks connRunnerCallbacks, + queueControlFrame func(wire.Frame), + generator ConnectionIDGenerator, +) *connIDGenerator { + m := &connIDGenerator{ + generator: generator, + activeSrcConnIDs: make(map[uint64]protocol.ConnectionID), + statelessResetter: statelessResetter, + connRunners: map[connRunner]connRunnerCallbacks{runner: callbacks}, + queueControlFrame: queueControlFrame, + } + m.activeSrcConnIDs[0] = initialConnectionID + m.initialClientDestConnID = initialClientDestConnID + return m +} + +func (m *connIDGenerator) SetMaxActiveConnIDs(limit uint64) error { + if m.generator.ConnectionIDLen() == 0 { + return nil + } + // The active_connection_id_limit transport parameter is the number of + // connection IDs the peer will store. This limit includes the connection ID + // used during the handshake, and the one sent in the preferred_address + // transport parameter. + // We currently don't send the preferred_address transport parameter, + // so we can issue (limit - 1) connection IDs. + for i := uint64(len(m.activeSrcConnIDs)); i < min(limit, protocol.MaxIssuedConnectionIDs); i++ { + if err := m.issueNewConnID(); err != nil { + return err + } + } + return nil +} + +func (m *connIDGenerator) Retire(seq uint64, sentWithDestConnID protocol.ConnectionID, expiry monotime.Time) error { + if seq > m.highestSeq { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: fmt.Sprintf("retired connection ID %d (highest issued: %d)", seq, m.highestSeq), + } + } + connID, ok := m.activeSrcConnIDs[seq] + // We might already have deleted this connection ID, if this is a duplicate frame. + if !ok { + return nil + } + if connID == sentWithDestConnID { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: fmt.Sprintf("retired connection ID %d (%s), which was used as the Destination Connection ID on this packet", seq, connID), + } + } + m.queueConnIDForRetiring(connID, expiry) + + delete(m.activeSrcConnIDs, seq) + // Don't issue a replacement for the initial connection ID. + if seq == 0 { + return nil + } + return m.issueNewConnID() +} + +func (m *connIDGenerator) queueConnIDForRetiring(connID protocol.ConnectionID, expiry monotime.Time) { + idx := slices.IndexFunc(m.connIDsToRetire, func(c connIDToRetire) bool { + return c.t.After(expiry) + }) + if idx == -1 { + idx = len(m.connIDsToRetire) + } + m.connIDsToRetire = slices.Insert(m.connIDsToRetire, idx, connIDToRetire{t: expiry, connID: connID}) +} + +func (m *connIDGenerator) issueNewConnID() error { + connID, err := m.generator.GenerateConnectionID() + if err != nil { + return err + } + m.activeSrcConnIDs[m.highestSeq+1] = connID + m.connRunners.AddConnectionID(connID) + m.queueControlFrame(&wire.NewConnectionIDFrame{ + SequenceNumber: m.highestSeq + 1, + ConnectionID: connID, + StatelessResetToken: m.statelessResetter.GetStatelessResetToken(connID), + }) + m.highestSeq++ + return nil +} + +func (m *connIDGenerator) SetHandshakeComplete(connIDExpiry monotime.Time) { + if m.initialClientDestConnID != nil { + m.queueConnIDForRetiring(*m.initialClientDestConnID, connIDExpiry) + m.initialClientDestConnID = nil + } +} + +func (m *connIDGenerator) RemoveRetiredConnIDs(now monotime.Time) { + if len(m.connIDsToRetire) == 0 { + return + } + for _, c := range m.connIDsToRetire { + if c.t.After(now) { + break + } + m.connRunners.RemoveConnectionID(c.connID) + m.connIDsToRetire = m.connIDsToRetire[1:] + } +} + +func (m *connIDGenerator) RemoveAll() { + if m.initialClientDestConnID != nil { + m.connRunners.RemoveConnectionID(*m.initialClientDestConnID) + } + for _, connID := range m.activeSrcConnIDs { + m.connRunners.RemoveConnectionID(connID) + } + for _, c := range m.connIDsToRetire { + m.connRunners.RemoveConnectionID(c.connID) + } +} + +func (m *connIDGenerator) ReplaceWithClosed(connClose []byte, expiry time.Duration) { + connIDs := make([]protocol.ConnectionID, 0, len(m.activeSrcConnIDs)+len(m.connIDsToRetire)+1) + if m.initialClientDestConnID != nil { + connIDs = append(connIDs, *m.initialClientDestConnID) + } + for _, connID := range m.activeSrcConnIDs { + connIDs = append(connIDs, connID) + } + for _, c := range m.connIDsToRetire { + connIDs = append(connIDs, c.connID) + } + m.connRunners.ReplaceWithClosed(connIDs, connClose, expiry) +} + +func (m *connIDGenerator) AddConnRunner(runner connRunner, r connRunnerCallbacks) { + // The transport might have already been added earlier. + // This happens if the application migrates back to and old path. + if _, ok := m.connRunners[runner]; ok { + return + } + m.connRunners[runner] = r + if m.initialClientDestConnID != nil { + r.AddConnectionID(*m.initialClientDestConnID) + } + for _, connID := range m.activeSrcConnIDs { + r.AddConnectionID(connID) + } +} diff --git a/third_party/quic-go/conn_id_generator_test.go b/third_party/quic-go/conn_id_generator_test.go new file mode 100644 index 0000000..869ba26 --- /dev/null +++ b/third_party/quic-go/conn_id_generator_test.go @@ -0,0 +1,350 @@ +package quic + +import ( + "math/rand/v2" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestConnIDGeneratorIssueAndRetire(t *testing.T) { + t.Run("with initial client destination connection ID", func(t *testing.T) { + testConnIDGeneratorIssueAndRetire(t, true) + }) + t.Run("without initial client destination connection ID", func(t *testing.T) { + testConnIDGeneratorIssueAndRetire(t, false) + }) +} + +func testConnIDGeneratorIssueAndRetire(t *testing.T, hasInitialClientDestConnID bool) { + var ( + added []protocol.ConnectionID + removed []protocol.ConnectionID + ) + var queuedFrames []wire.Frame + sr := newStatelessResetter(&StatelessResetKey{1, 2, 3, 4}) + var initialClientDestConnID *protocol.ConnectionID + if hasInitialClientDestConnID { + connID := protocol.ParseConnectionID([]byte{2, 2, 2, 2}) + initialClientDestConnID = &connID + } + g := newConnIDGenerator( + &packetHandlerMap{}, + protocol.ParseConnectionID([]byte{1, 1, 1, 1}), + initialClientDestConnID, + sr, + connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { added = append(added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { removed = append(removed, c) }, + ReplaceWithClosed: func([]protocol.ConnectionID, []byte, time.Duration) {}, + }, + func(f wire.Frame) { queuedFrames = append(queuedFrames, f) }, + &protocol.DefaultConnectionIDGenerator{ConnLen: 5}, + ) + + require.Empty(t, added) + require.NoError(t, g.SetMaxActiveConnIDs(4)) + require.Len(t, added, 3) + require.Len(t, queuedFrames, 3) + require.Empty(t, removed) + connIDs := make(map[uint64]protocol.ConnectionID) + // connection IDs 1, 2 and 3 were issued + for i, f := range queuedFrames { + ncid := f.(*wire.NewConnectionIDFrame) + require.EqualValues(t, i+1, ncid.SequenceNumber) + require.Equal(t, ncid.ConnectionID, added[i]) + require.Equal(t, ncid.StatelessResetToken, sr.GetStatelessResetToken(ncid.ConnectionID)) + connIDs[ncid.SequenceNumber] = ncid.ConnectionID + } + + // completing the handshake retires the initial client destination connection ID + added = added[:0] + queuedFrames = queuedFrames[:0] + now := monotime.Now() + g.SetHandshakeComplete(now) + require.Empty(t, added) + require.Empty(t, queuedFrames) + require.Empty(t, removed) + g.RemoveRetiredConnIDs(now) + if hasInitialClientDestConnID { + require.Equal(t, []protocol.ConnectionID{*initialClientDestConnID}, removed) + removed = removed[:0] + } else { + require.Empty(t, removed) + } + + // it's invalid to retire a connection ID that hasn't been issued yet + err := g.Retire(4, protocol.ParseConnectionID([]byte{3, 3, 3, 3}), monotime.Now()) + require.ErrorIs(t, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}, err) + require.ErrorContains(t, err, "retired connection ID 4 (highest issued: 3)") + // it's invalid to retire a connection ID in a packet that uses that connection ID + err = g.Retire(3, connIDs[3], monotime.Now()) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + require.ErrorContains(t, err, "was used as the Destination Connection ID on this packet") + + // retiring a connection ID makes us issue a new one + require.NoError(t, g.Retire(2, protocol.ParseConnectionID([]byte{3, 3, 3, 3}), monotime.Now())) + g.RemoveRetiredConnIDs(monotime.Now()) + require.Equal(t, []protocol.ConnectionID{connIDs[2]}, removed) + require.Len(t, queuedFrames, 1) + require.EqualValues(t, 4, queuedFrames[0].(*wire.NewConnectionIDFrame).SequenceNumber) + queuedFrames = queuedFrames[:0] + removed = removed[:0] + + // duplicate retirements don't do anything + require.NoError(t, g.Retire(2, protocol.ParseConnectionID([]byte{3, 3, 3, 3}), monotime.Now())) + g.RemoveRetiredConnIDs(monotime.Now()) + require.Empty(t, queuedFrames) + require.Empty(t, removed) +} + +func TestConnIDGeneratorRetiring(t *testing.T) { + initialConnID := protocol.ParseConnectionID([]byte{2, 2, 2, 2}) + var added, removed []protocol.ConnectionID + g := newConnIDGenerator( + &packetHandlerMap{}, + protocol.ParseConnectionID([]byte{1, 1, 1, 1}), + &initialConnID, + newStatelessResetter(&StatelessResetKey{1, 2, 3, 4}), + connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { added = append(added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { removed = append(removed, c) }, + ReplaceWithClosed: func([]protocol.ConnectionID, []byte, time.Duration) {}, + }, + func(f wire.Frame) {}, + &protocol.DefaultConnectionIDGenerator{ConnLen: 5}, + ) + require.NoError(t, g.SetMaxActiveConnIDs(6)) + require.Empty(t, removed) + require.Len(t, added, 5) + + now := monotime.Now() + + retirements := map[protocol.ConnectionID]monotime.Time{} + t1 := now.Add(time.Duration(rand.IntN(1000)) * time.Millisecond) + retirements[initialConnID] = t1 + g.SetHandshakeComplete(t1) + for i := range 5 { + t2 := now.Add(time.Duration(rand.IntN(1000)) * time.Millisecond) + require.NoError(t, g.Retire(uint64(i+1), protocol.ParseConnectionID([]byte{9, 9, 9, 9}), t2)) + retirements[added[i]] = t2 + + if rand.IntN(2) == 0 { + now = now.Add(time.Duration(rand.IntN(500)) * time.Millisecond) + g.RemoveRetiredConnIDs(now) + for _, r := range removed { + require.Contains(t, retirements, r) + require.LessOrEqual(t, retirements[r], now) + delete(retirements, r) + } + removed = removed[:0] + for _, r := range retirements { + require.Greater(t, r, now) + } + } + } +} + +func TestConnIDGeneratorRemoveAll(t *testing.T) { + t.Run("with initial client destination connection ID", func(t *testing.T) { + testConnIDGeneratorRemoveAll(t, true) + }) + t.Run("without initial client destination connection ID", func(t *testing.T) { + testConnIDGeneratorRemoveAll(t, false) + }) +} + +func testConnIDGeneratorRemoveAll(t *testing.T, hasInitialClientDestConnID bool) { + var initialClientDestConnID *protocol.ConnectionID + if hasInitialClientDestConnID { + connID := protocol.ParseConnectionID([]byte{2, 2, 2, 2}) + initialClientDestConnID = &connID + } + var ( + added []protocol.ConnectionID + removed []protocol.ConnectionID + ) + g := newConnIDGenerator( + &packetHandlerMap{}, + protocol.ParseConnectionID([]byte{1, 1, 1, 1}), + initialClientDestConnID, + newStatelessResetter(&StatelessResetKey{1, 2, 3, 4}), + connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { added = append(added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { removed = append(removed, c) }, + ReplaceWithClosed: func([]protocol.ConnectionID, []byte, time.Duration) {}, + }, + func(f wire.Frame) {}, + &protocol.DefaultConnectionIDGenerator{ConnLen: 5}, + ) + + require.NoError(t, g.SetMaxActiveConnIDs(1000)) + require.Len(t, added, protocol.MaxIssuedConnectionIDs-1) + + g.RemoveAll() + if hasInitialClientDestConnID { + require.Len(t, removed, protocol.MaxIssuedConnectionIDs+1) + require.Contains(t, removed, *initialClientDestConnID) + } else { + require.Len(t, removed, protocol.MaxIssuedConnectionIDs) + } + for _, id := range added { + require.Contains(t, removed, id) + } + require.Contains(t, removed, protocol.ParseConnectionID([]byte{1, 1, 1, 1})) +} + +func TestConnIDGeneratorReplaceWithClosed(t *testing.T) { + t.Run("with initial client destination connection ID", func(t *testing.T) { + testConnIDGeneratorReplaceWithClosed(t, true) + }) + t.Run("without initial client destination connection ID", func(t *testing.T) { + testConnIDGeneratorReplaceWithClosed(t, false) + }) +} + +func testConnIDGeneratorReplaceWithClosed(t *testing.T, hasInitialClientDestConnID bool) { + var initialClientDestConnID *protocol.ConnectionID + if hasInitialClientDestConnID { + connID := protocol.ParseConnectionID([]byte{2, 2, 2, 2}) + initialClientDestConnID = &connID + } + var ( + added []protocol.ConnectionID + replaced []protocol.ConnectionID + replacedWith []byte + ) + g := newConnIDGenerator( + &packetHandlerMap{}, + protocol.ParseConnectionID([]byte{1, 1, 1, 1}), + initialClientDestConnID, + newStatelessResetter(&StatelessResetKey{1, 2, 3, 4}), + connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { added = append(added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { t.Fatal("didn't expect conn ID removals") }, + ReplaceWithClosed: func(connIDs []protocol.ConnectionID, b []byte, _ time.Duration) { + replaced = connIDs + replacedWith = b + }, + }, + func(f wire.Frame) {}, + &protocol.DefaultConnectionIDGenerator{ConnLen: 5}, + ) + + require.NoError(t, g.SetMaxActiveConnIDs(1000)) + require.Len(t, added, protocol.MaxIssuedConnectionIDs-1) + // Retire two of these connection ID. + // This makes us issue two more connection IDs. + require.NoError(t, g.Retire(3, protocol.ParseConnectionID([]byte{1, 1, 1, 1}), monotime.Now())) + require.NoError(t, g.Retire(4, protocol.ParseConnectionID([]byte{1, 1, 1, 1}), monotime.Now())) + require.Len(t, added, protocol.MaxIssuedConnectionIDs+1) + + g.ReplaceWithClosed([]byte("foobar"), time.Second) + if hasInitialClientDestConnID { + require.Len(t, replaced, protocol.MaxIssuedConnectionIDs+3) + require.Contains(t, replaced, *initialClientDestConnID) + } else { + require.Len(t, replaced, protocol.MaxIssuedConnectionIDs+2) + } + for _, id := range added { + require.Contains(t, replaced, id) + } + require.Contains(t, replaced, protocol.ParseConnectionID([]byte{1, 1, 1, 1})) + require.Equal(t, []byte("foobar"), replacedWith) +} + +func TestConnIDGeneratorAddConnRunner(t *testing.T) { + initialConnID := protocol.ParseConnectionID([]byte{1, 1, 1, 1}) + clientDestConnID := protocol.ParseConnectionID([]byte{2, 2, 2, 2}) + + type connIDTracker struct { + added, removed, replaced []protocol.ConnectionID + } + + var tracker1, tracker2, tracker3 connIDTracker + runner1 := connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { tracker1.added = append(tracker1.added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { tracker1.removed = append(tracker1.removed, c) }, + ReplaceWithClosed: func(connIDs []protocol.ConnectionID, _ []byte, _ time.Duration) { + tracker1.replaced = append(tracker1.replaced, connIDs...) + }, + } + runner2 := connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { tracker2.added = append(tracker2.added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { tracker2.removed = append(tracker2.removed, c) }, + ReplaceWithClosed: func(connIDs []protocol.ConnectionID, _ []byte, _ time.Duration) { + tracker2.replaced = append(tracker2.replaced, connIDs...) + }, + } + runner3 := connRunnerCallbacks{ + AddConnectionID: func(c protocol.ConnectionID) { tracker3.added = append(tracker3.added, c) }, + RemoveConnectionID: func(c protocol.ConnectionID) { tracker3.removed = append(tracker3.removed, c) }, + ReplaceWithClosed: func(connIDs []protocol.ConnectionID, _ []byte, _ time.Duration) { + tracker3.replaced = append(tracker3.replaced, connIDs...) + }, + } + + sr := newStatelessResetter(&StatelessResetKey{1, 2, 3, 4}) + var queuedFrames []wire.Frame + + tr := &packetHandlerMap{} + g := newConnIDGenerator( + tr, + initialConnID, + &clientDestConnID, + sr, + runner1, + func(f wire.Frame) { queuedFrames = append(queuedFrames, f) }, + &protocol.DefaultConnectionIDGenerator{ConnLen: 5}, + ) + require.NoError(t, g.SetMaxActiveConnIDs(3)) + require.Len(t, tracker1.added, 2) + + // add the second runner - it should get all existing connection IDs + g.AddConnRunner(&packetHandlerMap{}, runner2) + require.Len(t, tracker1.added, 2) // unchanged + require.Len(t, tracker2.added, 4) + require.Contains(t, tracker2.added, initialConnID) + require.Contains(t, tracker2.added, clientDestConnID) + require.Contains(t, tracker2.added, tracker1.added[0]) + require.Contains(t, tracker2.added, tracker1.added[1]) + + // adding the same transport again doesn't do anything + trCopy := tr + g.AddConnRunner(trCopy, runner3) + require.Empty(t, tracker3.added) + + var connIDToRetire protocol.ConnectionID + var seqToRetire uint64 + ncid := queuedFrames[0].(*wire.NewConnectionIDFrame) + connIDToRetire = ncid.ConnectionID + seqToRetire = ncid.SequenceNumber + + require.NoError(t, g.Retire(seqToRetire, protocol.ParseConnectionID([]byte{3, 3, 3, 3}), monotime.Now())) + g.RemoveRetiredConnIDs(monotime.Now()) + require.Equal(t, []protocol.ConnectionID{connIDToRetire}, tracker1.removed) + require.Equal(t, []protocol.ConnectionID{connIDToRetire}, tracker2.removed) + + tracker1.removed = nil + tracker2.removed = nil + g.SetHandshakeComplete(monotime.Now()) + g.RemoveRetiredConnIDs(monotime.Now()) + require.Equal(t, []protocol.ConnectionID{clientDestConnID}, tracker1.removed) + require.Equal(t, []protocol.ConnectionID{clientDestConnID}, tracker2.removed) + + g.ReplaceWithClosed([]byte("connection closed"), time.Second) + require.True(t, len(tracker1.replaced) > 0) + require.Equal(t, tracker1.replaced, tracker2.replaced) + + tracker1.removed = nil + tracker2.removed = nil + g.RemoveAll() + require.NotEmpty(t, tracker1.removed) + require.Equal(t, tracker1.removed, tracker2.removed) +} diff --git a/third_party/quic-go/conn_id_manager.go b/third_party/quic-go/conn_id_manager.go new file mode 100644 index 0000000..3b72fb3 --- /dev/null +++ b/third_party/quic-go/conn_id_manager.go @@ -0,0 +1,321 @@ +package quic + +import ( + "fmt" + "slices" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" +) + +type newConnID struct { + SequenceNumber uint64 + ConnectionID protocol.ConnectionID + StatelessResetToken protocol.StatelessResetToken +} + +type connIDManager struct { + queue []newConnID + + highestProbingID uint64 + pathProbing map[pathID]newConnID // initialized lazily + + handshakeComplete bool + activeSequenceNumber uint64 + highestRetired uint64 + activeConnectionID protocol.ConnectionID + activeStatelessResetToken *protocol.StatelessResetToken + + // We change the connection ID after sending on average + // protocol.PacketsPerConnectionID packets. The actual value is randomized + // hide the packet loss rate from on-path observers. + rand utils.Rand + packetsSinceLastChange uint32 + packetsPerConnectionID uint32 + + addStatelessResetToken func(protocol.StatelessResetToken) + removeStatelessResetToken func(protocol.StatelessResetToken) + queueControlFrame func(wire.Frame) + + closed bool +} + +func newConnIDManager( + initialDestConnID protocol.ConnectionID, + addStatelessResetToken func(protocol.StatelessResetToken), + removeStatelessResetToken func(protocol.StatelessResetToken), + queueControlFrame func(wire.Frame), +) *connIDManager { + return &connIDManager{ + activeConnectionID: initialDestConnID, + addStatelessResetToken: addStatelessResetToken, + removeStatelessResetToken: removeStatelessResetToken, + queueControlFrame: queueControlFrame, + queue: make([]newConnID, 0, protocol.MaxActiveConnectionIDs), + } +} + +func (h *connIDManager) AddFromPreferredAddress(connID protocol.ConnectionID, resetToken protocol.StatelessResetToken) error { + return h.addConnectionID(1, connID, resetToken) +} + +func (h *connIDManager) Add(f *wire.NewConnectionIDFrame) error { + if err := h.add(f); err != nil { + return err + } + if len(h.queue) >= protocol.MaxActiveConnectionIDs { + return &qerr.TransportError{ErrorCode: qerr.ConnectionIDLimitError} + } + return nil +} + +func (h *connIDManager) add(f *wire.NewConnectionIDFrame) error { + if h.activeConnectionID.Len() == 0 { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "received NEW_CONNECTION_ID frame but zero-length connection IDs are in use", + } + } + // If the NEW_CONNECTION_ID frame is reordered, such that its sequence number is smaller than the currently active + // connection ID or if it was already retired, send the RETIRE_CONNECTION_ID frame immediately. + if f.SequenceNumber < max(h.activeSequenceNumber, h.highestProbingID) || f.SequenceNumber < h.highestRetired { + h.queueControlFrame(&wire.RetireConnectionIDFrame{ + SequenceNumber: f.SequenceNumber, + }) + return nil + } + + if f.RetirePriorTo != 0 && h.pathProbing != nil { + for id, entry := range h.pathProbing { + if entry.SequenceNumber < f.RetirePriorTo { + h.queueControlFrame(&wire.RetireConnectionIDFrame{ + SequenceNumber: entry.SequenceNumber, + }) + h.removeStatelessResetToken(entry.StatelessResetToken) + delete(h.pathProbing, id) + } + } + } + // Retire elements in the queue. + // Doesn't retire the active connection ID. + if f.RetirePriorTo > h.highestRetired { + var newQueue []newConnID + for _, entry := range h.queue { + if entry.SequenceNumber >= f.RetirePriorTo { + newQueue = append(newQueue, entry) + } else { + h.queueControlFrame(&wire.RetireConnectionIDFrame{SequenceNumber: entry.SequenceNumber}) + } + } + h.queue = newQueue + h.highestRetired = f.RetirePriorTo + } + + if f.SequenceNumber == h.activeSequenceNumber { + return nil + } + + if err := h.addConnectionID(f.SequenceNumber, f.ConnectionID, f.StatelessResetToken); err != nil { + return err + } + + // Retire the active connection ID, if necessary. + if h.activeSequenceNumber < f.RetirePriorTo { + // The queue is guaranteed to have at least one element at this point. + h.updateConnectionID() + } + return nil +} + +func (h *connIDManager) addConnectionID(seq uint64, connID protocol.ConnectionID, resetToken protocol.StatelessResetToken) error { + // fast path: add to the end of the queue + if len(h.queue) == 0 || h.queue[len(h.queue)-1].SequenceNumber < seq { + h.queue = append(h.queue, newConnID{ + SequenceNumber: seq, + ConnectionID: connID, + StatelessResetToken: resetToken, + }) + return nil + } + + // slow path: insert in the middle + for i, entry := range h.queue { + if entry.SequenceNumber == seq { + if entry.ConnectionID != connID { + return fmt.Errorf("received conflicting connection IDs for sequence number %d", seq) + } + if entry.StatelessResetToken != resetToken { + return fmt.Errorf("received conflicting stateless reset tokens for sequence number %d", seq) + } + return nil + } + + // insert at the correct position to maintain sorted order + if entry.SequenceNumber > seq { + h.queue = slices.Insert(h.queue, i, newConnID{ + SequenceNumber: seq, + ConnectionID: connID, + StatelessResetToken: resetToken, + }) + return nil + } + } + return nil // unreachable +} + +func (h *connIDManager) updateConnectionID() { + h.assertNotClosed() + h.queueControlFrame(&wire.RetireConnectionIDFrame{ + SequenceNumber: h.activeSequenceNumber, + }) + h.highestRetired = max(h.highestRetired, h.activeSequenceNumber) + if h.activeStatelessResetToken != nil { + h.removeStatelessResetToken(*h.activeStatelessResetToken) + } + + front := h.queue[0] + h.queue = h.queue[1:] + h.activeSequenceNumber = front.SequenceNumber + h.activeConnectionID = front.ConnectionID + h.activeStatelessResetToken = &front.StatelessResetToken + h.packetsSinceLastChange = 0 + h.packetsPerConnectionID = protocol.PacketsPerConnectionID/2 + uint32(h.rand.Int31n(protocol.PacketsPerConnectionID)) + h.addStatelessResetToken(*h.activeStatelessResetToken) +} + +func (h *connIDManager) Close() { + h.closed = true + if h.activeStatelessResetToken != nil { + h.removeStatelessResetToken(*h.activeStatelessResetToken) + } + if h.pathProbing != nil { + for _, entry := range h.pathProbing { + h.removeStatelessResetToken(entry.StatelessResetToken) + } + } +} + +// is called when the server performs a Retry +// and when the server changes the connection ID in the first Initial sent +func (h *connIDManager) ChangeInitialConnID(newConnID protocol.ConnectionID) { + if h.activeSequenceNumber != 0 { + panic("expected first connection ID to have sequence number 0") + } + h.activeConnectionID = newConnID +} + +// is called when the server provides a stateless reset token in the transport parameters +func (h *connIDManager) SetStatelessResetToken(token protocol.StatelessResetToken) { + h.assertNotClosed() + if h.activeSequenceNumber != 0 { + panic("expected first connection ID to have sequence number 0") + } + h.activeStatelessResetToken = &token + h.addStatelessResetToken(token) +} + +func (h *connIDManager) SentPacket() { + h.packetsSinceLastChange++ +} + +func (h *connIDManager) shouldUpdateConnID() bool { + if !h.handshakeComplete { + return false + } + // initiate the first change as early as possible (after handshake completion) + if len(h.queue) > 0 && h.activeSequenceNumber == 0 { + return true + } + // For later changes, only change if + // 1. The queue of connection IDs is filled more than 50%. + // 2. We sent at least PacketsPerConnectionID packets + return 2*len(h.queue) >= protocol.MaxActiveConnectionIDs && + h.packetsSinceLastChange >= h.packetsPerConnectionID +} + +func (h *connIDManager) Get() protocol.ConnectionID { + h.assertNotClosed() + if h.shouldUpdateConnID() { + h.updateConnectionID() + } + return h.activeConnectionID +} + +func (h *connIDManager) SetHandshakeComplete() { + h.handshakeComplete = true +} + +// GetConnIDForPath retrieves a connection ID for a new path (i.e. not the active one). +// Once a connection ID is allocated for a path, it cannot be used for a different path. +// When called with the same pathID, it will return the same connection ID, +// unless the peer requested that this connection ID be retired. +func (h *connIDManager) GetConnIDForPath(id pathID) (protocol.ConnectionID, bool) { + h.assertNotClosed() + // if we're using zero-length connection IDs, we don't need to change the connection ID + if h.activeConnectionID.Len() == 0 { + return protocol.ConnectionID{}, true + } + + if h.pathProbing == nil { + h.pathProbing = make(map[pathID]newConnID) + } + entry, ok := h.pathProbing[id] + if ok { + return entry.ConnectionID, true + } + if len(h.queue) == 0 { + return protocol.ConnectionID{}, false + } + front := h.queue[0] + h.queue = h.queue[1:] + h.pathProbing[id] = front + h.highestProbingID = front.SequenceNumber + h.addStatelessResetToken(front.StatelessResetToken) + return front.ConnectionID, true +} + +func (h *connIDManager) RetireConnIDForPath(pathID pathID) { + h.assertNotClosed() + // if we're using zero-length connection IDs, we don't need to change the connection ID + if h.activeConnectionID.Len() == 0 { + return + } + + entry, ok := h.pathProbing[pathID] + if !ok { + return + } + h.queueControlFrame(&wire.RetireConnectionIDFrame{ + SequenceNumber: entry.SequenceNumber, + }) + h.removeStatelessResetToken(entry.StatelessResetToken) + delete(h.pathProbing, pathID) +} + +func (h *connIDManager) IsActiveStatelessResetToken(token protocol.StatelessResetToken) bool { + if h.activeStatelessResetToken != nil { + if *h.activeStatelessResetToken == token { + return true + } + } + if h.pathProbing != nil { + for _, entry := range h.pathProbing { + if entry.StatelessResetToken == token { + return true + } + } + } + return false +} + +// Using the connIDManager after it has been closed can have disastrous effects: +// If the connection ID is rotated, a new entry would be inserted into the packet handler map, +// leading to a memory leak of the connection struct. +// See https://github.com/apernet/quic-go/pull/4852 for more details. +func (h *connIDManager) assertNotClosed() { + if h.closed { + panic("connection ID manager is closed") + } +} diff --git a/third_party/quic-go/conn_id_manager_test.go b/third_party/quic-go/conn_id_manager_test.go new file mode 100644 index 0000000..00936c1 --- /dev/null +++ b/third_party/quic-go/conn_id_manager_test.go @@ -0,0 +1,432 @@ +package quic + +import ( + "crypto/rand" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestConnIDManagerInitialConnID(t *testing.T) { + m := newConnIDManager(protocol.ParseConnectionID([]byte{1, 2, 3, 4}), nil, nil, nil) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4}), m.Get()) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4}), m.Get()) + m.ChangeInitialConnID(protocol.ParseConnectionID([]byte{5, 6, 7, 8})) + require.Equal(t, protocol.ParseConnectionID([]byte{5, 6, 7, 8}), m.Get()) +} + +func TestConnIDManagerAddConnIDs(t *testing.T) { + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(protocol.StatelessResetToken) {}, + func(protocol.StatelessResetToken) {}, + func(wire.Frame) {}, + ) + f1 := &wire.NewConnectionIDFrame{ + SequenceNumber: 1, + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd, 0xe}, + } + f2 := &wire.NewConnectionIDFrame{ + SequenceNumber: 2, + ConnectionID: protocol.ParseConnectionID([]byte{0xba, 0xad, 0xf0, 0x0d}), + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd, 0xe}, + } + require.NoError(t, m.Add(f2)) + require.NoError(t, m.Add(f1)) // receiving reordered frames is fine + require.NoError(t, m.Add(f2)) // receiving a duplicate is fine + + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4}), m.Get()) + m.updateConnectionID() + require.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), m.Get()) + m.updateConnectionID() + require.Equal(t, protocol.ParseConnectionID([]byte{0xba, 0xad, 0xf0, 0x0d}), m.Get()) + + require.NoError(t, m.Add(f2)) // receiving a duplicate for the current connection ID is fine as well + require.Equal(t, protocol.ParseConnectionID([]byte{0xba, 0xad, 0xf0, 0x0d}), m.Get()) + + // receiving mismatching connection IDs is not fine + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 3, + ConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), // mismatching connection ID + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd, 0xe}, + })) + require.EqualError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 3, + ConnectionID: protocol.ParseConnectionID([]byte{2, 3, 4, 5}), // mismatching connection ID + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd, 0xe}, + }), "received conflicting connection IDs for sequence number 3") + // receiving mismatching stateless reset tokens is not fine either + require.EqualError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 3, + ConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + StatelessResetToken: protocol.StatelessResetToken{1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd, 0xe, 0}, + }), "received conflicting stateless reset tokens for sequence number 3") +} + +func TestConnIDManagerLimit(t *testing.T) { + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(protocol.StatelessResetToken) {}, + func(protocol.StatelessResetToken) {}, + func(f wire.Frame) {}, + ) + for i := uint8(1); i < protocol.MaxActiveConnectionIDs; i++ { + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: uint64(i), + ConnectionID: protocol.ParseConnectionID([]byte{i, i, i, i}), + StatelessResetToken: protocol.StatelessResetToken{i, i, i, i, i, i, i, i, i, i, i, i, i, i, i, i}, + })) + } + require.Equal(t, &qerr.TransportError{ErrorCode: qerr.ConnectionIDLimitError}, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: uint64(9999), + ConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + StatelessResetToken: protocol.StatelessResetToken{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}, + })) +} + +func TestConnIDManagerRetiringConnectionIDs(t *testing.T) { + var frameQueue []wire.Frame + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(protocol.StatelessResetToken) {}, + func(protocol.StatelessResetToken) {}, + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + ) + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 10, + ConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + })) + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 13, + ConnectionID: protocol.ParseConnectionID([]byte{2, 3, 4, 5}), + })) + require.Empty(t, frameQueue) + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + RetirePriorTo: 14, + SequenceNumber: 17, + ConnectionID: protocol.ParseConnectionID([]byte{3, 4, 5, 6}), + })) + require.Equal(t, []wire.Frame{ + &wire.RetireConnectionIDFrame{SequenceNumber: 10}, + &wire.RetireConnectionIDFrame{SequenceNumber: 13}, + &wire.RetireConnectionIDFrame{SequenceNumber: 0}, + }, frameQueue) + require.Equal(t, protocol.ParseConnectionID([]byte{3, 4, 5, 6}), m.Get()) + frameQueue = nil + + // a reordered connection ID is immediately retired + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 12, + ConnectionID: protocol.ParseConnectionID([]byte{5, 6, 7, 8}), + })) + require.Equal(t, []wire.Frame{&wire.RetireConnectionIDFrame{SequenceNumber: 12}}, frameQueue) + require.Equal(t, protocol.ParseConnectionID([]byte{3, 4, 5, 6}), m.Get()) +} + +func TestConnIDManagerHandshakeCompletion(t *testing.T) { + var frameQueue []wire.Frame + var addedTokens, removedTokens []protocol.StatelessResetToken + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(token protocol.StatelessResetToken) { addedTokens = append(addedTokens, token) }, + func(token protocol.StatelessResetToken) { removedTokens = append(removedTokens, token) }, + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + ) + m.SetStatelessResetToken(protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}) + require.Equal(t, []protocol.StatelessResetToken{{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}}, addedTokens) + require.Empty(t, removedTokens) + + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 1, + ConnectionID: protocol.ParseConnectionID([]byte{4, 3, 2, 1}), + StatelessResetToken: protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}, + })) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4}), m.Get()) + m.SetHandshakeComplete() + require.Equal(t, protocol.ParseConnectionID([]byte{4, 3, 2, 1}), m.Get()) + require.Equal(t, []wire.Frame{&wire.RetireConnectionIDFrame{SequenceNumber: 0}}, frameQueue) + require.Equal(t, []protocol.StatelessResetToken{{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}}, removedTokens) +} + +func TestConnIDManagerConnIDRotation(t *testing.T) { + toToken := func(connID protocol.ConnectionID) protocol.StatelessResetToken { + var token protocol.StatelessResetToken + copy(token[:], connID.Bytes()) + copy(token[connID.Len():], connID.Bytes()) + return token + } + + var frameQueue []wire.Frame + var addedTokens, removedTokens []protocol.StatelessResetToken + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(token protocol.StatelessResetToken) { addedTokens = append(addedTokens, token) }, + func(token protocol.StatelessResetToken) { removedTokens = append(removedTokens, token) }, + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + ) + // the first connection ID is used as soon as the handshake is complete + m.SetHandshakeComplete() + firstConnID := protocol.ParseConnectionID([]byte{4, 3, 2, 1}) + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 1, + ConnectionID: firstConnID, + StatelessResetToken: toToken(protocol.ParseConnectionID([]byte{4, 3, 2, 1})), + })) + require.Equal(t, firstConnID, m.Get()) + frameQueue = nil + require.True(t, m.IsActiveStatelessResetToken(toToken(firstConnID))) + require.Equal(t, addedTokens, []protocol.StatelessResetToken{toToken(firstConnID)}) + addedTokens = addedTokens[:0] + + // Note that we're missing the connection ID with sequence number 2. + // It will be received later. + var queuedConnIDs []protocol.ConnectionID + for i := range protocol.MaxActiveConnectionIDs - 1 { + b := make([]byte, 4) + rand.Read(b) + connID := protocol.ParseConnectionID(b) + queuedConnIDs = append(queuedConnIDs, connID) + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: uint64(3 + i), + ConnectionID: connID, + StatelessResetToken: toToken(connID), + })) + require.False(t, m.IsActiveStatelessResetToken(toToken(connID))) + } + + var counter int + for { + require.Empty(t, frameQueue) + m.SentPacket() + counter++ + if connID := m.Get(); connID != firstConnID { + require.Equal(t, queuedConnIDs[0], m.Get()) + require.Equal(t, []wire.Frame{&wire.RetireConnectionIDFrame{SequenceNumber: 1}}, frameQueue) + require.Equal(t, removedTokens, []protocol.StatelessResetToken{toToken(firstConnID)}) + require.Equal(t, addedTokens, []protocol.StatelessResetToken{toToken(connID)}) + addedTokens = addedTokens[:0] + removedTokens = removedTokens[:0] + require.True(t, m.IsActiveStatelessResetToken(toToken(connID))) + require.False(t, m.IsActiveStatelessResetToken(toToken(firstConnID))) + break + } + require.True(t, m.IsActiveStatelessResetToken(toToken(firstConnID))) + require.Empty(t, addedTokens) + } + require.GreaterOrEqual(t, counter, protocol.PacketsPerConnectionID/2) + require.LessOrEqual(t, counter, protocol.PacketsPerConnectionID*3/2) + frameQueue = nil + + // now receive connection ID 2 + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 2, + ConnectionID: protocol.ParseConnectionID([]byte{2, 3, 4, 5}), + })) + require.Equal(t, []wire.Frame{&wire.RetireConnectionIDFrame{SequenceNumber: 2}}, frameQueue) +} + +func TestConnIDManagerPathMigration(t *testing.T) { + var frameQueue []wire.Frame + var addedTokens, removedTokens []protocol.StatelessResetToken + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(token protocol.StatelessResetToken) { addedTokens = append(addedTokens, token) }, + func(token protocol.StatelessResetToken) { removedTokens = append(removedTokens, token) }, + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + ) + + // no connection ID available yet + _, ok := m.GetConnIDForPath(1) + require.False(t, ok) + + // add two connection IDs + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 1, + ConnectionID: protocol.ParseConnectionID([]byte{4, 3, 2, 1}), + StatelessResetToken: protocol.StatelessResetToken{4, 3, 2, 1, 4, 3, 2, 1}, + })) + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 2, + ConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2}), + StatelessResetToken: protocol.StatelessResetToken{5, 4, 3, 2, 5, 4, 3, 2}, + })) + connID, ok := m.GetConnIDForPath(1) + require.True(t, ok) + require.Equal(t, protocol.ParseConnectionID([]byte{4, 3, 2, 1}), connID) + require.Equal(t, []protocol.StatelessResetToken{{4, 3, 2, 1, 4, 3, 2, 1}}, addedTokens) + require.Empty(t, removedTokens) + + addedTokens = addedTokens[:0] + require.False(t, m.IsActiveStatelessResetToken(protocol.StatelessResetToken{5, 4, 3, 2, 5, 4, 3, 2})) + connID, ok = m.GetConnIDForPath(2) + require.True(t, ok) + require.Equal(t, protocol.ParseConnectionID([]byte{5, 4, 3, 2}), connID) + require.Equal(t, []protocol.StatelessResetToken{{5, 4, 3, 2, 5, 4, 3, 2}}, addedTokens) + require.Empty(t, removedTokens) + require.True(t, m.IsActiveStatelessResetToken(protocol.StatelessResetToken{5, 4, 3, 2, 5, 4, 3, 2})) + + addedTokens = addedTokens[:0] + // asking for the connection for path 1 again returns the same connection ID + connID, ok = m.GetConnIDForPath(1) + require.True(t, ok) + require.Equal(t, protocol.ParseConnectionID([]byte{4, 3, 2, 1}), connID) + require.Empty(t, addedTokens) + + // if the connection ID is retired, the path will use another connection ID + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 3, + RetirePriorTo: 2, + ConnectionID: protocol.ParseConnectionID([]byte{6, 5, 4, 3}), + StatelessResetToken: protocol.StatelessResetToken{6, 5, 4, 3, 6, 5, 4, 3}, + })) + require.Len(t, frameQueue, 2) + require.Equal(t, []protocol.StatelessResetToken{{4, 3, 2, 1, 4, 3, 2, 1}}, removedTokens) + frameQueue = nil + removedTokens = removedTokens[:0] + + require.Equal(t, protocol.ParseConnectionID([]byte{6, 5, 4, 3}), m.Get()) + require.Equal(t, []protocol.StatelessResetToken{{6, 5, 4, 3, 6, 5, 4, 3}}, addedTokens) + require.Empty(t, removedTokens) + addedTokens = addedTokens[:0] + + // the connection ID is not used for new paths + _, ok = m.GetConnIDForPath(3) + require.False(t, ok) + + // Manually retiring the connection ID does nothing. + // Path 1 doesn't have a connection ID anymore. + m.RetireConnIDForPath(1) + require.Empty(t, frameQueue) + _, ok = m.GetConnIDForPath(1) + require.False(t, ok) + require.Empty(t, removedTokens) + + // only after a new connection ID is added, it will be used for path 1 + require.NoError(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 4, + ConnectionID: protocol.ParseConnectionID([]byte{7, 6, 5, 4}), + StatelessResetToken: protocol.StatelessResetToken{16, 15, 14, 13}, + })) + connID, ok = m.GetConnIDForPath(1) + require.True(t, ok) + require.Equal(t, protocol.ParseConnectionID([]byte{7, 6, 5, 4}), connID) + require.Equal(t, []protocol.StatelessResetToken{{16, 15, 14, 13}}, addedTokens) + require.Empty(t, removedTokens) + require.True(t, m.IsActiveStatelessResetToken(protocol.StatelessResetToken{16, 15, 14, 13})) + + // a RETIRE_CONNECTION_ID frame for path 1 is queued when retiring the connection ID + m.RetireConnIDForPath(1) + require.Equal(t, []wire.Frame{&wire.RetireConnectionIDFrame{SequenceNumber: 4}}, frameQueue) + require.Equal(t, []protocol.StatelessResetToken{{16, 15, 14, 13}}, removedTokens) + removedTokens = removedTokens[:0] + require.False(t, m.IsActiveStatelessResetToken(protocol.StatelessResetToken{16, 15, 14, 13})) + + m.Close() + require.Equal(t, []protocol.StatelessResetToken{ + {6, 5, 4, 3, 6, 5, 4, 3}, // currently active connection ID + {5, 4, 3, 2, 5, 4, 3, 2}, // path 2 + }, removedTokens) +} + +func TestConnIDManagerZeroLengthConnectionID(t *testing.T) { + m := newConnIDManager( + protocol.ConnectionID{}, + func(protocol.StatelessResetToken) {}, + func(protocol.StatelessResetToken) {}, + func(f wire.Frame) {}, + ) + require.Equal(t, protocol.ConnectionID{}, m.Get()) + for range 5 * protocol.PacketsPerConnectionID { + m.SentPacket() + require.Equal(t, protocol.ConnectionID{}, m.Get()) + } + + // for path probing, we don't need to change the connection ID + for id := pathID(1); id < 10; id++ { + connID, ok := m.GetConnIDForPath(id) + require.True(t, ok) + require.Equal(t, protocol.ConnectionID{}, connID) + } + // retiring a connection ID for a path is also a no-op + for id := pathID(1); id < 20; id++ { + m.RetireConnIDForPath(id) + } + + require.ErrorIs(t, m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: 1, + ConnectionID: protocol.ConnectionID{}, + StatelessResetToken: protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}, + }), &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) +} + +func TestConnIDManagerClose(t *testing.T) { + var addedTokens, removedTokens []protocol.StatelessResetToken + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(token protocol.StatelessResetToken) { addedTokens = append(addedTokens, token) }, + func(token protocol.StatelessResetToken) { removedTokens = append(removedTokens, token) }, + func(f wire.Frame) {}, + ) + m.SetStatelessResetToken(protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}) + require.Equal(t, []protocol.StatelessResetToken{{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}}, addedTokens) + require.Empty(t, removedTokens) + m.Close() + require.Equal(t, []protocol.StatelessResetToken{{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}}, removedTokens) + + require.Panics(t, func() { m.Get() }) + require.Panics(t, func() { m.SetStatelessResetToken(protocol.StatelessResetToken{}) }) +} + +func BenchmarkConnIDManagerReordered(b *testing.B) { + benchmarkConnIDManager(b, true) +} + +func BenchmarkConnIDManagerInOrder(b *testing.B) { + benchmarkConnIDManager(b, false) +} + +func benchmarkConnIDManager(b *testing.B, reordered bool) { + m := newConnIDManager( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + func(protocol.StatelessResetToken) {}, + func(protocol.StatelessResetToken) {}, + func(f wire.Frame) {}, + ) + connIDs := make([]protocol.ConnectionID, 0, protocol.MaxActiveConnectionIDs) + statelessResetTokens := make([]protocol.StatelessResetToken, 0, protocol.MaxActiveConnectionIDs) + for range protocol.MaxActiveConnectionIDs { + b := make([]byte, 8) + rand.Read(b) + connIDs = append(connIDs, protocol.ParseConnectionID(b)) + var statelessResetToken protocol.StatelessResetToken + rand.Read(statelessResetToken[:]) + statelessResetTokens = append(statelessResetTokens, statelessResetToken) + } + + // 1 -> 3 + // 2 -> 1 + // 3 -> 2 + // 4 -> 4 + offsets := []int{2, -1, -1, 0} + + b.ResetTimer() + for i := range b.N { + seq := i + if reordered { + seq += offsets[i%len(offsets)] + } + m.Add(&wire.NewConnectionIDFrame{ + SequenceNumber: uint64(seq), + ConnectionID: connIDs[i%len(connIDs)], + StatelessResetToken: statelessResetTokens[i%len(statelessResetTokens)], + }) + if i > protocol.MaxActiveConnectionIDs-2 { + m.updateConnectionID() + } + } +} diff --git a/third_party/quic-go/conn_wrapped_test.go b/third_party/quic-go/conn_wrapped_test.go new file mode 100644 index 0000000..6dd8d81 --- /dev/null +++ b/third_party/quic-go/conn_wrapped_test.go @@ -0,0 +1,73 @@ +package quic + +import "context" + +func (c *wrappedConn) run() error { + if c.testHooks == nil { + return c.Conn.run() + } + if c.testHooks.run != nil { + return c.testHooks.run() + } + return nil +} + +func (c *wrappedConn) earlyConnReady() <-chan struct{} { + if c.testHooks == nil { + return c.Conn.earlyConnReady() + } + if c.testHooks.earlyConnReady != nil { + return c.testHooks.earlyConnReady() + } + return nil +} + +func (c *wrappedConn) Context() context.Context { + if c.testHooks == nil { + return c.Conn.Context() + } + if c.testHooks.context != nil { + return c.testHooks.context() + } + return context.Background() +} + +func (c *wrappedConn) HandshakeComplete() <-chan struct{} { + if c.testHooks == nil { + return c.Conn.HandshakeComplete() + } + if c.testHooks.handshakeComplete != nil { + return c.testHooks.handshakeComplete() + } + return nil +} + +func (c *wrappedConn) closeWithTransportError(code TransportErrorCode) { + if c.testHooks == nil { + c.Conn.closeWithTransportError(code) + return + } + if c.testHooks.closeWithTransportError != nil { + c.testHooks.closeWithTransportError(code) + } +} + +func (c *wrappedConn) destroy(e error) { + if c.testHooks == nil { + c.Conn.destroy(e) + return + } + if c.testHooks.destroy != nil { + c.testHooks.destroy(e) + } +} + +func (c *wrappedConn) handlePacket(p receivedPacket) { + if c.testHooks == nil { + c.Conn.handlePacket(p) + return + } + if c.testHooks.handlePacket != nil { + c.testHooks.handlePacket(p) + } +} diff --git a/third_party/quic-go/connection.go b/third_party/quic-go/connection.go new file mode 100644 index 0000000..0a71b0b --- /dev/null +++ b/third_party/quic-go/connection.go @@ -0,0 +1,3214 @@ +package quic + +import ( + "bytes" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "reflect" + "slices" + "sync" + "sync/atomic" + "time" + + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/utils/ringbuffer" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +type unpacker interface { + UnpackLongHeader(hdr *wire.Header, data []byte) (*unpackedPacket, error) + UnpackShortHeader(rcvTime monotime.Time, data []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error) +} + +type cryptoStreamHandler interface { + StartHandshake(context.Context) error + ChangeConnectionID(protocol.ConnectionID) + SetLargest1RTTAcked(protocol.PacketNumber) error + SetHandshakeConfirmed() + GetSessionTicket() ([]byte, error) + NextEvent() handshake.Event + DiscardInitialKeys() + HandleMessage([]byte, protocol.EncryptionLevel) error + io.Closer + ConnectionState() handshake.ConnectionState +} + +type receivedPacket struct { + buffer *packetBuffer + + remoteAddr net.Addr + rcvTime monotime.Time + data []byte + + ecn protocol.ECN + + info packetInfo // only valid if the contained IP address is valid +} + +type receivedPacketWithChecksum struct { + receivedPacket + checksum qlog.DatagramPayloadChecksum +} + +func (p *receivedPacket) Size() protocol.ByteCount { return protocol.ByteCount(len(p.data)) } + +func (p *receivedPacket) Clone() *receivedPacket { + return &receivedPacket{ + remoteAddr: p.remoteAddr, + rcvTime: p.rcvTime, + data: p.data, + buffer: p.buffer, + ecn: p.ecn, + info: p.info, + } +} + +type connRunner interface { + Add(protocol.ConnectionID, packetHandler) bool + Remove(protocol.ConnectionID) + ReplaceWithClosed([]protocol.ConnectionID, []byte, time.Duration) + AddResetToken(protocol.StatelessResetToken, packetHandler) + RemoveResetToken(protocol.StatelessResetToken) +} + +type closeError struct { + err error + immediate bool +} + +type errCloseForRecreating struct { + nextPacketNumber protocol.PacketNumber + nextVersion protocol.Version +} + +func (e *errCloseForRecreating) Error() string { + return "closing connection in order to recreate it" +} + +var deadlineSendImmediately = monotime.Time(42 * time.Millisecond) // any value > time.Time{} and before time.Now() is fine + +type blockMode uint8 + +const ( + // blockModeNone means that the connection is not blocked. + blockModeNone blockMode = iota + // blockModeCongestionLimited means that the connection is congestion limited. + // In that case, we can still send acknowledgments and PTO probe packets. + blockModeCongestionLimited + // blockModeHardBlocked means that no packet can be sent, under no circumstances. This can happen when: + // * the send queue is full + // * the SentPacketHandler returns SendNone, e.g. when we are tracking the maximum number of packets + // In that case, the timer will be set to the idle timeout. + blockModeHardBlocked +) + +// A Conn is a QUIC connection between two peers. +// Calls to the connection (and to streams) can return the following types of errors: +// - [ApplicationError]: for errors triggered by the application running on top of QUIC +// - [TransportError]: for errors triggered by the QUIC transport (in many cases a misbehaving peer) +// - [IdleTimeoutError]: when the peer goes away unexpectedly (this is a [net.Error] timeout error) +// - [HandshakeTimeoutError]: when the cryptographic handshake takes too long (this is a [net.Error] timeout error) +// - [StatelessResetError]: when we receive a stateless reset +// - [VersionNegotiationError]: returned by the client, when there's no version overlap between the peers +type Conn struct { + // Destination connection ID used during the handshake. + // Used to check source connection ID on incoming packets. + handshakeDestConnID protocol.ConnectionID + // Set for the client. Destination connection ID used on the first Initial sent. + origDestConnID protocol.ConnectionID + retrySrcConnID *protocol.ConnectionID // only set for the client (and if a Retry was performed) + + srcConnIDLen int + + perspective protocol.Perspective + version protocol.Version + config *Config + + conn sendConn + sendQueue sender + + // lazily initialzed: most connections never migrate + pathManager *pathManager + largestRcvdAppData protocol.PacketNumber + pathManagerOutgoing atomic.Pointer[pathManagerOutgoing] + + streamsMap *streamsMap + connIDManager *connIDManager + connIDGenerator *connIDGenerator + + rttStats *utils.RTTStats + connStats utils.ConnectionStats + + cryptoStreamManager *cryptoStreamManager + sentPacketHandler ackhandler.SentPacketHandler + receivedPacketHandler ackhandler.ReceivedPacketHandler + retransmissionQueue *retransmissionQueue + framer *framer + connFlowController *connectionFlowController + tokenStoreKey string // only set for the client + tokenGenerator *handshake.TokenGenerator // only set for the server + + unpacker unpacker + frameParser wire.FrameParser + packer packer + mtuDiscoverer mtuDiscoverer // initialized when the transport parameters are received + + maxPayloadSizeEstimate atomic.Uint32 + + initialStream *initialCryptoStream + handshakeStream *cryptoStream + oneRTTStream *cryptoStream // only set for the server + cryptoStreamHandler cryptoStreamHandler + + notifyReceivedPacket chan struct{} + sendingScheduled chan struct{} + receivedPacketMx sync.Mutex + receivedPackets ringbuffer.RingBuffer[receivedPacket] + + // closeChan is used to notify the run loop that it should terminate + closeChan chan struct{} + closeErr atomic.Pointer[closeError] + + ctx context.Context + ctxCancel context.CancelCauseFunc + handshakeCompleteChan chan struct{} + + undecryptablePackets []receivedPacketWithChecksum // undecryptable packets, waiting for a change in encryption level + undecryptablePacketsToProcess []receivedPacketWithChecksum + + earlyConnReadyChan chan struct{} + sentFirstPacket bool + droppedInitialKeys bool + handshakeComplete bool + handshakeConfirmed bool + + receivedRetry bool + versionNegotiated bool + receivedFirstPacket bool + + blocked blockMode + + // the minimum of the max_idle_timeout values advertised by both endpoints + idleTimeout time.Duration + creationTime monotime.Time + // The idle timeout is set based on the max of the time we received the last packet... + lastPacketReceivedTime monotime.Time + // ... and the time we sent a new ack-eliciting packet after receiving a packet. + firstAckElicitingPacketAfterIdleSentTime monotime.Time + // pacingDeadline is the time when the next packet should be sent + pacingDeadline monotime.Time + + peerParams *wire.TransportParameters + + timer *time.Timer + // keepAlivePingSent stores whether a keep alive PING is in flight. + // It is reset as soon as we receive a packet from the peer. + keepAlivePingSent bool + keepAliveInterval time.Duration + + datagramQueue *datagramQueue + + connStateMutex sync.Mutex + connState ConnectionState + + logID string + qlogTrace qlogwriter.Trace + qlogger qlogwriter.Recorder + logger utils.Logger +} + +var _ streamSender = &Conn{} + +type connTestHooks struct { + run func() error + earlyConnReady func() <-chan struct{} + context func() context.Context + handshakeComplete func() <-chan struct{} + closeWithTransportError func(TransportErrorCode) + destroy func(error) + handlePacket func(receivedPacket) +} + +type wrappedConn struct { + testHooks *connTestHooks + *Conn +} + +var newConnection = func( + ctx context.Context, + ctxCancel context.CancelCauseFunc, + conn sendConn, + runner connRunner, + origDestConnID protocol.ConnectionID, + retrySrcConnID *protocol.ConnectionID, + clientDestConnID protocol.ConnectionID, + destConnID protocol.ConnectionID, + srcConnID protocol.ConnectionID, + connIDGenerator ConnectionIDGenerator, + statelessResetter *statelessResetter, + conf *Config, + tlsConf *tls.Config, + tokenGenerator *handshake.TokenGenerator, + clientAddressValidated bool, + rtt time.Duration, + qlogTrace qlogwriter.Trace, + logger utils.Logger, + v protocol.Version, +) *wrappedConn { + s := &Conn{ + ctx: ctx, + ctxCancel: ctxCancel, + conn: conn, + config: conf, + handshakeDestConnID: destConnID, + srcConnIDLen: srcConnID.Len(), + tokenGenerator: tokenGenerator, + oneRTTStream: newCryptoStream(), + perspective: protocol.PerspectiveServer, + qlogTrace: qlogTrace, + logger: logger, + version: v, + } + if qlogTrace != nil { + s.qlogger = qlogTrace.AddProducer() + } + if origDestConnID.Len() > 0 { + s.logID = origDestConnID.String() + } else { + s.logID = destConnID.String() + } + s.connIDManager = newConnIDManager( + destConnID, + func(token protocol.StatelessResetToken) { runner.AddResetToken(token, s) }, + runner.RemoveResetToken, + s.queueControlFrame, + ) + s.connIDGenerator = newConnIDGenerator( + runner, + srcConnID, + &clientDestConnID, + statelessResetter, + connRunnerCallbacks{ + AddConnectionID: func(connID protocol.ConnectionID) { runner.Add(connID, s) }, + RemoveConnectionID: runner.Remove, + ReplaceWithClosed: runner.ReplaceWithClosed, + }, + s.queueControlFrame, + connIDGenerator, + ) + s.preSetup() + s.rttStats.SetInitialRTT(rtt) + s.sentPacketHandler = ackhandler.NewSentPacketHandler( + 0, + protocol.ByteCount(s.config.InitialPacketSize), + s.rttStats, + &s.connStats, + clientAddressValidated, + s.conn.capabilities().ECN, + s.receivedPacketHandler.IgnorePacketsBelow, + s.perspective, + false, // servers keep quic-go's 2-byte packet number floor + s.qlogger, + s.logger, + ) + s.maxPayloadSizeEstimate.Store(uint32(estimateMaxPayloadSize(protocol.ByteCount(s.config.InitialPacketSize)))) + statelessResetToken := statelessResetter.GetStatelessResetToken(srcConnID) + params := &wire.TransportParameters{ + InitialMaxStreamDataBidiLocal: protocol.ByteCount(s.config.InitialStreamReceiveWindow), + InitialMaxStreamDataBidiRemote: protocol.ByteCount(s.config.InitialStreamReceiveWindow), + InitialMaxStreamDataUni: protocol.ByteCount(s.config.InitialStreamReceiveWindow), + InitialMaxData: protocol.ByteCount(s.config.InitialConnectionReceiveWindow), + MaxIdleTimeout: s.config.MaxIdleTimeout, + MaxBidiStreamNum: protocol.StreamNum(s.config.MaxIncomingStreams), + MaxUniStreamNum: protocol.StreamNum(s.config.MaxIncomingUniStreams), + MaxAckDelay: protocol.MaxAckDelayInclGranularity, + AckDelayExponent: protocol.AckDelayExponent, + MaxUDPPayloadSize: protocol.MaxPacketBufferSize, + StatelessResetToken: &statelessResetToken, + OriginalDestinationConnectionID: origDestConnID, + // For interoperability with quic-go versions before May 2023, this value must be set to a value + // different from protocol.DefaultActiveConnectionIDLimit. + // If set to the default value, it will be omitted from the transport parameters, which will make + // old quic-go versions interpret it as 0, instead of the default value of 2. + // See https://github.com/apernet/quic-go/pull/3806. + ActiveConnectionIDLimit: protocol.MaxActiveConnectionIDs, + InitialSourceConnectionID: srcConnID, + RetrySourceConnectionID: retrySrcConnID, + EnableResetStreamAt: conf.EnableStreamResetPartialDelivery, + } + if s.config.EnableDatagrams && !s.config.OmitMaxDatagramFrameSize { + params.MaxDatagramFrameSize = wire.MaxDatagramSize + if s.config.MaxDatagramFrameSize != 0 { + params.MaxDatagramFrameSize = protocol.ByteCount(s.config.MaxDatagramFrameSize) + } + } else { + params.MaxDatagramFrameSize = protocol.InvalidByteCount + } + if s.qlogger != nil { + s.qlogTransportParameters(params, protocol.PerspectiveServer, false) + } + cs := handshake.NewCryptoSetupServer( + clientDestConnID, + conn.LocalAddr(), + conn.RemoteAddr(), + params, + tlsConf, + conf.Allow0RTT, + s.rttStats, + s.qlogger, + logger, + s.version, + ) + s.cryptoStreamHandler = cs + s.packer = newPacketPacker(srcConnID, s.connIDManager.Get, s.initialStream, s.handshakeStream, s.sentPacketHandler, s.retransmissionQueue, cs, s.framer, &s.receivedPacketHandler, s.datagramQueue, s.perspective, false) + s.unpacker = newPacketUnpacker(cs, s.srcConnIDLen) + s.cryptoStreamManager = newCryptoStreamManager(s.initialStream, s.handshakeStream, s.oneRTTStream) + return &wrappedConn{Conn: s} +} + +// declare this as a variable, such that we can it mock it in the tests +var newClientConnection = func( + ctx context.Context, + conn sendConn, + runner connRunner, + destConnID protocol.ConnectionID, + srcConnID protocol.ConnectionID, + connIDGenerator ConnectionIDGenerator, + statelessResetter *statelessResetter, + conf *Config, + tlsConf *tls.Config, + initialPacketNumber protocol.PacketNumber, + enable0RTT bool, + hasNegotiatedVersion bool, + qlogTrace qlogwriter.Trace, + logger utils.Logger, + v protocol.Version, +) (*wrappedConn, error) { + s := &Conn{ + conn: conn, + config: conf, + origDestConnID: destConnID, + handshakeDestConnID: destConnID, + srcConnIDLen: srcConnID.Len(), + perspective: protocol.PerspectiveClient, + logID: destConnID.String(), + logger: logger, + qlogTrace: qlogTrace, + versionNegotiated: hasNegotiatedVersion, + version: v, + } + if qlogTrace != nil { + s.qlogger = qlogTrace.AddProducer() + } + if s.qlogger != nil { + var srcAddr, destAddr *net.UDPAddr + if addr, ok := conn.LocalAddr().(*net.UDPAddr); ok { + srcAddr = addr + } + if addr, ok := conn.RemoteAddr().(*net.UDPAddr); ok { + destAddr = addr + } + s.qlogger.RecordEvent(startedConnectionEvent(srcAddr, destAddr)) + } + s.connIDManager = newConnIDManager( + destConnID, + func(token protocol.StatelessResetToken) { runner.AddResetToken(token, s) }, + runner.RemoveResetToken, + s.queueControlFrame, + ) + s.connIDGenerator = newConnIDGenerator( + runner, + srcConnID, + nil, + statelessResetter, + connRunnerCallbacks{ + AddConnectionID: func(connID protocol.ConnectionID) { runner.Add(connID, s) }, + RemoveConnectionID: runner.Remove, + ReplaceWithClosed: runner.ReplaceWithClosed, + }, + s.queueControlFrame, + connIDGenerator, + ) + s.ctx, s.ctxCancel = context.WithCancelCause(ctx) + s.preSetup() + s.sentPacketHandler = ackhandler.NewSentPacketHandler( + initialPacketNumber, + protocol.ByteCount(s.config.InitialPacketSize), + s.rttStats, + &s.connStats, + false, // has no effect + s.conn.capabilities().ECN, + s.receivedPacketHandler.IgnorePacketsBelow, + s.perspective, + s.config.ChromeParrot, + s.qlogger, + s.logger, + ) + s.maxPayloadSizeEstimate.Store(uint32(estimateMaxPayloadSize(protocol.ByteCount(s.config.InitialPacketSize)))) + oneRTTStream := newCryptoStream() + params := &wire.TransportParameters{ + InitialMaxStreamDataBidiRemote: protocol.ByteCount(s.config.InitialStreamReceiveWindow), + InitialMaxStreamDataBidiLocal: protocol.ByteCount(s.config.InitialStreamReceiveWindow), + InitialMaxStreamDataUni: protocol.ByteCount(s.config.InitialStreamReceiveWindow), + InitialMaxData: protocol.ByteCount(s.config.InitialConnectionReceiveWindow), + MaxIdleTimeout: s.config.MaxIdleTimeout, + MaxBidiStreamNum: protocol.StreamNum(s.config.MaxIncomingStreams), + MaxUniStreamNum: protocol.StreamNum(s.config.MaxIncomingUniStreams), + MaxAckDelay: protocol.MaxAckDelayInclGranularity, + MaxUDPPayloadSize: protocol.MaxPacketBufferSize, + AckDelayExponent: protocol.AckDelayExponent, + // For interoperability with quic-go versions before May 2023, this value must be set to a value + // different from protocol.DefaultActiveConnectionIDLimit. + // If set to the default value, it will be omitted from the transport parameters, which will make + // old quic-go versions interpret it as 0, instead of the default value of 2. + // See https://github.com/apernet/quic-go/pull/3806. + ActiveConnectionIDLimit: protocol.MaxActiveConnectionIDs, + InitialSourceConnectionID: srcConnID, + EnableResetStreamAt: conf.EnableStreamResetPartialDelivery, + } + if s.config.EnableDatagrams && !s.config.OmitMaxDatagramFrameSize { + params.MaxDatagramFrameSize = wire.MaxDatagramSize + if s.config.MaxDatagramFrameSize != 0 { + params.MaxDatagramFrameSize = protocol.ByteCount(s.config.MaxDatagramFrameSize) + } + } else { + params.MaxDatagramFrameSize = protocol.InvalidByteCount + } + if s.config.ChromeParrot { + chromeParrotTransportParameters(params) + } + if s.qlogger != nil { + s.qlogTransportParameters(params, protocol.PerspectiveClient, false) + } + cs, err := handshake.NewCryptoSetupClient( + destConnID, + params, + tlsConf, + enable0RTT, + s.config.ChromeParrot, + s.rttStats, + s.qlogger, + logger, + s.version, + ) + if err != nil { + return nil, err + } + s.cryptoStreamHandler = cs + s.cryptoStreamManager = newCryptoStreamManager(s.initialStream, s.handshakeStream, oneRTTStream) + s.unpacker = newPacketUnpacker(cs, s.srcConnIDLen) + s.packer = newPacketPacker(srcConnID, s.connIDManager.Get, s.initialStream, s.handshakeStream, s.sentPacketHandler, s.retransmissionQueue, cs, s.framer, &s.receivedPacketHandler, s.datagramQueue, s.perspective, s.config.ChromeParrot) + if len(tlsConf.ServerName) > 0 { + s.tokenStoreKey = tlsConf.ServerName + } else { + s.tokenStoreKey = conn.RemoteAddr().String() + } + if s.config.TokenStore != nil { + if token := s.config.TokenStore.Pop(s.tokenStoreKey); token != nil { + s.packer.SetToken(token.data) + s.rttStats.SetInitialRTT(token.rtt) + } + } + return &wrappedConn{Conn: s}, nil +} + +func (c *Conn) preSetup() { + c.largestRcvdAppData = protocol.InvalidPacketNumber + c.initialStream = newInitialCryptoStream(c.perspective == protocol.PerspectiveClient, c.config.ChromeParrot) + c.handshakeStream = newCryptoStream() + c.sendQueue = newSendQueue(c.conn) + c.retransmissionQueue = newRetransmissionQueue() + c.frameParser = *wire.NewFrameParser( + c.config.EnableDatagrams, + c.config.EnableStreamResetPartialDelivery, + false, // ACK_FREQUENCY is not supported yet + ) + c.rttStats = utils.NewRTTStats() + c.connFlowController = newConnectionFlowController( + protocol.ByteCount(c.config.InitialConnectionReceiveWindow), + protocol.ByteCount(c.config.MaxConnectionReceiveWindow), + func(size protocol.ByteCount) bool { + if c.config.AllowConnectionWindowIncrease == nil { + return true + } + return c.config.AllowConnectionWindowIncrease(c, uint64(size)) + }, + c.rttStats, + c.logger, + ) + c.earlyConnReadyChan = make(chan struct{}) + c.streamsMap = newStreamsMap( + c.ctx, + c, + c.queueControlFrame, + c.newFlowController, + uint64(c.config.MaxIncomingStreams), + uint64(c.config.MaxIncomingUniStreams), + c.perspective, + ) + c.framer = newFramer(c.connFlowController) + c.receivedPackets.Init(8) + c.notifyReceivedPacket = make(chan struct{}, 1) + c.closeChan = make(chan struct{}, 1) + c.sendingScheduled = make(chan struct{}, 1) + c.handshakeCompleteChan = make(chan struct{}) + + now := monotime.Now() + c.lastPacketReceivedTime = now + c.creationTime = now + + c.receivedPacketHandler = *ackhandler.NewReceivedPacketHandler(c.logger) + + c.datagramQueue = newDatagramQueue(c.scheduleSending, c.logger) + c.connState.Version = c.version +} + +// run the connection main loop +func (c *Conn) run() (err error) { + defer func() { c.ctxCancel(err) }() + + defer func() { + // drain queued packets that will never be processed + c.receivedPacketMx.Lock() + defer c.receivedPacketMx.Unlock() + + for !c.receivedPackets.Empty() { + p := c.receivedPackets.PopFront() + p.buffer.Decrement() + p.buffer.MaybeRelease() + } + }() + + c.timer = time.NewTimer(monotime.Until(c.idleTimeoutStartTime().Add(c.config.HandshakeIdleTimeout))) + + if err := c.cryptoStreamHandler.StartHandshake(c.ctx); err != nil { + return err + } + if err := c.handleHandshakeEvents(monotime.Now()); err != nil { + return err + } + go func() { + if err := c.sendQueue.Run(); err != nil { + c.destroyImpl(err) + } + }() + + if c.perspective == protocol.PerspectiveClient { + c.scheduleSending() // so the ClientHello actually gets sent + } + + var sendQueueAvailable <-chan struct{} + +runLoop: + for { + if c.framer.QueuedTooManyControlFrames() { + c.setCloseError(&closeError{err: &qerr.TransportError{ErrorCode: InternalError}}) + break runLoop + } + // Close immediately if requested + select { + case <-c.closeChan: + break runLoop + default: + } + + // no need to set a timer if we can send packets immediately + if c.pacingDeadline != deadlineSendImmediately { + c.maybeResetTimer() + } + + // 1st: handle undecryptable packets, if any. + // This can only occur before completion of the handshake. + if len(c.undecryptablePacketsToProcess) > 0 { + var processedUndecryptablePacket bool + queue := c.undecryptablePacketsToProcess + c.undecryptablePacketsToProcess = nil + for _, p := range queue { + processed, err := c.handleOnePacket(p.receivedPacket, p.checksum) + if err != nil { + c.setCloseError(&closeError{err: err}) + break runLoop + } + if processed { + processedUndecryptablePacket = true + } + } + if processedUndecryptablePacket { + // if we processed any undecryptable packets, jump to the resetting of the timers directly + continue + } + } + + // 2nd: receive packets. + processed, err := c.handlePackets() // don't check receivedPackets.Len() in the run loop to avoid locking the mutex + if err != nil { + c.setCloseError(&closeError{err: err}) + break runLoop + } + + // We don't need to wait for new events if: + // * we processed packets: we probably need to send an ACK, and potentially more data + // * the pacer allows us to send more packets immediately + shouldProceedImmediately := sendQueueAvailable == nil && (processed || c.pacingDeadline.Equal(deadlineSendImmediately)) + if !shouldProceedImmediately { + // 3rd: wait for something to happen: + // * closing of the connection + // * timer firing + // * sending scheduled + // * send queue available + // * received packets + select { + case <-c.closeChan: + break runLoop + case <-c.timer.C: + case <-c.sendingScheduled: + case <-sendQueueAvailable: + case <-c.notifyReceivedPacket: + wasProcessed, err := c.handlePackets() + if err != nil { + c.setCloseError(&closeError{err: err}) + break runLoop + } + // if we processed any undecryptable packets, jump to the resetting of the timers directly + if !wasProcessed { + continue + } + } + } + + // Check for loss detection timeout. + // This could cause packets to be declared lost, and retransmissions to be enqueued. + now := monotime.Now() + if timeout := c.sentPacketHandler.GetLossDetectionTimeout(); !timeout.IsZero() && !timeout.After(now) { + if err := c.sentPacketHandler.OnLossDetectionTimeout(now); err != nil { + c.setCloseError(&closeError{err: err}) + break runLoop + } + } + + if keepAliveTime := c.nextKeepAliveTime(); !keepAliveTime.IsZero() && !now.Before(keepAliveTime) { + // send a PING frame since there is no activity in the connection + c.logger.Debugf("Sending a keep-alive PING to keep the connection alive.") + c.framer.QueueControlFrame(&wire.PingFrame{}) + c.keepAlivePingSent = true + } else if !c.handshakeComplete && now.Sub(c.creationTime) >= c.config.handshakeTimeout() { + c.destroyImpl(qerr.ErrHandshakeTimeout) + break runLoop + } else { + idleTimeoutStartTime := c.idleTimeoutStartTime() + if (!c.handshakeComplete && now.Sub(idleTimeoutStartTime) >= c.config.HandshakeIdleTimeout) || + (c.handshakeComplete && !now.Before(c.nextIdleTimeoutTime())) { + c.destroyImpl(qerr.ErrIdleTimeout) + break runLoop + } + } + + c.connIDGenerator.RemoveRetiredConnIDs(now) + + if c.perspective == protocol.PerspectiveClient { + pm := c.pathManagerOutgoing.Load() + if pm != nil { + tr, ok := pm.ShouldSwitchPath() + if ok { + c.switchToNewPath(tr, now) + } + } + } + + if c.sendQueue.WouldBlock() { + // The send queue is still busy sending out packets. Wait until there's space to enqueue new packets. + sendQueueAvailable = c.sendQueue.Available() + // Cancel the pacing timer, as we can't send any more packets until the send queue is available again. + c.pacingDeadline = 0 + c.blocked = blockModeHardBlocked + continue + } + + if c.closeErr.Load() != nil { + break runLoop + } + + c.blocked = blockModeNone // sending might set it back to true if we're congestion limited + if err := c.triggerSending(now); err != nil { + c.setCloseError(&closeError{err: err}) + break runLoop + } + if c.sendQueue.WouldBlock() { + // The send queue is still busy sending out packets. Wait until there's space to enqueue new packets. + sendQueueAvailable = c.sendQueue.Available() + // Cancel the pacing timer, as we can't send any more packets until the send queue is available again. + c.pacingDeadline = 0 + c.blocked = blockModeHardBlocked + } else { + sendQueueAvailable = nil + } + } + + closeErr := c.closeErr.Load() + c.cryptoStreamHandler.Close() + c.sendQueue.Close() // close the send queue before sending the CONNECTION_CLOSE + c.handleCloseError(closeErr) + if c.qlogger != nil { + if e := (&errCloseForRecreating{}); !errors.As(closeErr.err, &e) { + c.qlogger.Close() + } + } + c.logger.Infof("Connection %s closed.", c.logID) + c.timer.Stop() + return closeErr.err +} + +// blocks until the early connection can be used +func (c *Conn) earlyConnReady() <-chan struct{} { + return c.earlyConnReadyChan +} + +// Context returns a context that is cancelled when the connection is closed. +// The cancellation cause is set to the error that caused the connection to close. +func (c *Conn) Context() context.Context { + return c.ctx +} + +func (c *Conn) supportsDatagrams() bool { + return c.peerMaxDatagramFrameSize() > 0 +} + +func (c *Conn) peerMaxDatagramFrameSize() protocol.ByteCount { + if c.peerParams.MaxDatagramFrameSize > 0 { + return c.peerParams.MaxDatagramFrameSize + } + if c.config.EnableDatagrams && c.config.AssumePeerMaxDatagramFrameSize > 0 { + return protocol.ByteCount(c.config.AssumePeerMaxDatagramFrameSize) + } + return protocol.InvalidByteCount +} + +// ConnectionState returns basic details about the QUIC connection. +func (c *Conn) ConnectionState() ConnectionState { + c.connStateMutex.Lock() + defer c.connStateMutex.Unlock() + + cs := c.cryptoStreamHandler.ConnectionState() + c.connState.TLS = cs.ConnectionState + c.connState.Used0RTT = cs.Used0RTT + if c.peerParams != nil { + c.connState.SupportsDatagrams.Remote = c.supportsDatagrams() + c.connState.SupportsStreamResetPartialDelivery.Remote = c.peerParams.EnableResetStreamAt + } + c.connState.SupportsDatagrams.Local = c.config.EnableDatagrams + c.connState.SupportsStreamResetPartialDelivery.Local = c.config.EnableStreamResetPartialDelivery + c.connState.GSO = c.conn.capabilities().GSO + return c.connState +} + +// ConnectionStats contains statistics about the QUIC connection +type ConnectionStats struct { + // MinRTT is the estimate of the minimum RTT observed on the active network + // path. + MinRTT time.Duration + // LatestRTT is the last RTT sample observed on the active network path. + LatestRTT time.Duration + // SmoothedRTT is an exponentially weighted moving average of an endpoint's + // RTT samples. See https://www.rfc-editor.org/rfc/rfc9002#section-5.3 + SmoothedRTT time.Duration + // MeanDeviation estimates the variation in the RTT samples using a mean + // variation. See https://www.rfc-editor.org/rfc/rfc9002#section-5.3 + MeanDeviation time.Duration + + // BytesSent is the number of bytes sent on the underlying connection, + // including retransmissions. Does not include UDP or any other outer + // framing. + BytesSent uint64 + // PacketsSent is the number of packets sent on the underlying connection, + // including those that are determined to have been lost. + PacketsSent uint64 + // BytesReceived is the number of total bytes received on the underlying + // connection, including duplicate data for streams. Does not include UDP or + // any other outer framing. + BytesReceived uint64 + // PacketsReceived is the number of total packets received on the underlying + // connection, including packets that were not processable. + PacketsReceived uint64 + // BytesLost is the number of bytes lost on the underlying connection (does + // not monotonically increase, because packets that are declared lost can + // subsequently be received). Does not include UDP or any other outer + // framing. + BytesLost uint64 + // PacketsLost is the number of packets lost on the underlying connection + // (does not monotonically increase, because packets that are declared lost + // can subsequently be received). + PacketsLost uint64 +} + +func (c *Conn) ConnectionStats() ConnectionStats { + return ConnectionStats{ + MinRTT: c.rttStats.MinRTT(), + LatestRTT: c.rttStats.LatestRTT(), + SmoothedRTT: c.rttStats.SmoothedRTT(), + MeanDeviation: c.rttStats.MeanDeviation(), + + BytesSent: c.connStats.BytesSent.Load(), + PacketsSent: c.connStats.PacketsSent.Load(), + BytesReceived: c.connStats.BytesReceived.Load(), + PacketsReceived: c.connStats.PacketsReceived.Load(), + BytesLost: c.connStats.BytesLost.Load(), + PacketsLost: c.connStats.PacketsLost.Load(), + } +} + +// Time when the connection should time out +func (c *Conn) nextIdleTimeoutTime() monotime.Time { + idleTimeout := max(c.idleTimeout, c.rttStats.PTO(true)*3) + return c.idleTimeoutStartTime().Add(idleTimeout) +} + +// Time when the next keep-alive packet should be sent. +// It returns a zero time if no keep-alive should be sent. +func (c *Conn) nextKeepAliveTime() monotime.Time { + if c.config.KeepAlivePeriod == 0 || c.keepAlivePingSent { + return 0 + } + keepAliveInterval := max(c.keepAliveInterval, c.rttStats.PTO(true)*3/2) + return c.lastPacketReceivedTime.Add(keepAliveInterval) +} + +func (c *Conn) maybeResetTimer() { + var deadline monotime.Time + if !c.handshakeComplete { + deadline = c.creationTime.Add(c.config.handshakeTimeout()) + if t := c.idleTimeoutStartTime().Add(c.config.HandshakeIdleTimeout); t.Before(deadline) { + deadline = t + } + } else { + // A keep-alive packet is ack-eliciting, so it can only be sent if the connection is + // neither congestion limited nor hard-blocked. + if c.blocked != blockModeNone { + deadline = c.nextIdleTimeoutTime() + } else { + if keepAliveTime := c.nextKeepAliveTime(); !keepAliveTime.IsZero() { + deadline = keepAliveTime + } else { + deadline = c.nextIdleTimeoutTime() + } + } + } + // If the connection is hard-blocked, we can't even send acknowledgments, + // nor can we send PTO probe packets. + if c.blocked == blockModeHardBlocked { + c.timer.Reset(monotime.Until(deadline)) + return + } + + if t := c.receivedPacketHandler.GetAlarmTimeout(); !t.IsZero() && t.Before(deadline) { + deadline = t + } + if t := c.sentPacketHandler.GetLossDetectionTimeout(); !t.IsZero() && t.Before(deadline) { + deadline = t + } + if c.blocked == blockModeCongestionLimited { + c.timer.Reset(monotime.Until(deadline)) + return + } + + if !c.pacingDeadline.IsZero() && c.pacingDeadline.Before(deadline) { + deadline = c.pacingDeadline + } + c.timer.Reset(monotime.Until(deadline)) +} + +func (c *Conn) idleTimeoutStartTime() monotime.Time { + startTime := c.lastPacketReceivedTime + if t := c.firstAckElicitingPacketAfterIdleSentTime; !t.IsZero() && t.After(startTime) { + startTime = t + } + return startTime +} + +func (c *Conn) switchToNewPath(tr *Transport, now monotime.Time) { + initialPacketSize := protocol.ByteCount(c.config.InitialPacketSize) + c.sentPacketHandler.MigratedPath(now, initialPacketSize) + maxPacketSize := protocol.ByteCount(protocol.MaxPacketBufferSize) + if c.peerParams.MaxUDPPayloadSize > 0 && c.peerParams.MaxUDPPayloadSize < maxPacketSize { + maxPacketSize = c.peerParams.MaxUDPPayloadSize + } + c.mtuDiscoverer.Reset(now, initialPacketSize, maxPacketSize) + c.conn = newSendConn(tr.conn, c.conn.RemoteAddr(), packetInfo{}, utils.DefaultLogger) // TODO: find a better way + c.sendQueue.Close() + c.sendQueue = newSendQueue(c.conn) + go func() { + if err := c.sendQueue.Run(); err != nil { + c.destroyImpl(err) + } + }() +} + +func (c *Conn) handleHandshakeComplete(now monotime.Time) error { + defer close(c.handshakeCompleteChan) + // Once the handshake completes, we have derived 1-RTT keys. + // There's no point in queueing undecryptable packets for later decryption anymore. + c.undecryptablePackets = nil + + c.connIDManager.SetHandshakeComplete() + c.connIDGenerator.SetHandshakeComplete(now.Add(3 * c.rttStats.PTO(false))) + + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.ALPNInformation{ + ChosenALPN: c.cryptoStreamHandler.ConnectionState().NegotiatedProtocol, + }) + } + + // The server applies transport parameters right away, but the client side has to wait for handshake completion. + // During a 0-RTT connection, the client is only allowed to use the new transport parameters for 1-RTT packets. + if c.perspective == protocol.PerspectiveClient { + c.applyTransportParameters() + return nil + } + + // All these only apply to the server side. + if err := c.handleHandshakeConfirmed(now); err != nil { + return err + } + + ticket, err := c.cryptoStreamHandler.GetSessionTicket() + if err != nil { + return err + } + if ticket != nil { // may be nil if session tickets are disabled via tls.Config.SessionTicketsDisabled + c.oneRTTStream.Write(ticket) + for c.oneRTTStream.HasData() { + if cf := c.oneRTTStream.PopCryptoFrame(protocol.MaxPostHandshakeCryptoFrameSize); cf != nil { + c.queueControlFrame(cf) + } + } + } + token, err := c.tokenGenerator.NewToken(c.conn.RemoteAddr(), c.rttStats.SmoothedRTT()) + if err != nil { + return err + } + c.queueControlFrame(&wire.NewTokenFrame{Token: token}) + c.queueControlFrame(&wire.HandshakeDoneFrame{}) + return nil +} + +func (c *Conn) handleHandshakeConfirmed(now monotime.Time) error { + // Drop initial keys. + // On the client side, this should have happened when sending the first Handshake packet, + // but this is not guaranteed if the server misbehaves. + // See CVE-2025-59530 for more details. + if err := c.dropEncryptionLevel(protocol.EncryptionInitial, now); err != nil { + return err + } + if err := c.dropEncryptionLevel(protocol.EncryptionHandshake, now); err != nil { + return err + } + + c.handshakeConfirmed = true + c.cryptoStreamHandler.SetHandshakeConfirmed() + + if !c.config.DisablePathMTUDiscovery && c.conn.capabilities().DF { + c.mtuDiscoverer.Start(now) + } + return nil +} + +const maxPacketsToProcess = 32 + +func (c *Conn) handlePackets() (wasProcessed bool, _ error) { + // Process packets from the receivedPackets queue. + // Limit the number of packets to process to maxPacketsToProcess, + // so we eventually get a chance to send out an ACK when receiving a lot of packets. + c.receivedPacketMx.Lock() + + if c.receivedPackets.Empty() { + c.receivedPacketMx.Unlock() + return false, nil + } + + var hasMorePackets bool + for range maxPacketsToProcess { + p := c.receivedPackets.PopFront() + c.receivedPacketMx.Unlock() + + var datagramPayloadChecksum qlog.DatagramPayloadChecksum + if c.qlogger != nil && wire.IsLongHeaderPacket(p.data[0]) { + datagramPayloadChecksum = qlog.CalculateDatagramPayloadChecksum(p.data) + } + processed, err := c.handleOnePacket(p, datagramPayloadChecksum) + if err != nil { + return false, err + } + if processed { + wasProcessed = true + } + c.receivedPacketMx.Lock() + hasMorePackets = !c.receivedPackets.Empty() + if !hasMorePackets { + break + } + // Prioritize sending of new CRYPTO data. + // This is especially relevant when processing 0-RTT packets. + if !c.handshakeComplete && (c.initialStream.HasData() || c.handshakeStream.HasData()) { + break + } + } + c.receivedPacketMx.Unlock() + + if hasMorePackets { + select { + case c.notifyReceivedPacket <- struct{}{}: + default: + } + } + return wasProcessed, nil +} + +func (c *Conn) handleOnePacket(rp receivedPacket, datagramPayloadChecksum qlog.DatagramPayloadChecksum) (wasProcessed bool, _ error) { + c.sentPacketHandler.ReceivedBytes(rp.Size(), rp.rcvTime) + + if wire.IsVersionNegotiationPacket(rp.data) { + return false, c.handleVersionNegotiationPacket(rp) + } + + var counter uint8 + var lastConnID protocol.ConnectionID + data := rp.data + p := rp + for len(data) > 0 { + if counter > 0 { + p = *(p.Clone()) + p.data = data + + destConnID, err := wire.ParseConnectionID(p.data, c.srcConnIDLen) + if err != nil { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: len(data)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + c.logger.Debugf("error parsing packet, couldn't parse connection ID: %s", err) + break + } + if destConnID != lastConnID { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{DestConnectionID: destConnID}, + Raw: qlog.RawInfo{Length: len(data)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropUnknownConnectionID, + }) + } + c.logger.Debugf("coalesced packet has different destination connection ID: %s, expected %s", destConnID, lastConnID) + break + } + } + + if wire.IsLongHeaderPacket(p.data[0]) { + hdr, packetData, rest, err := wire.ParsePacket(p.data) + if err != nil { + if c.qlogger != nil { + if err == wire.ErrUnsupportedVersion { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{Version: hdr.Version}, + Raw: qlog.RawInfo{Length: len(data)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropUnsupportedVersion, + }) + } else { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: len(data)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + } + c.logger.Debugf("error parsing packet: %s", err) + break + } + lastConnID = hdr.DestConnectionID + + if hdr.Version != c.version { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: len(data)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropUnexpectedVersion, + }) + } + c.logger.Debugf("Dropping packet with version %x. Expected %x.", hdr.Version, c.version) + break + } + + if counter > 0 { + p.buffer.Split() + } + counter++ + + // only log if this actually a coalesced packet + if c.logger.Debug() && (counter > 1 || len(rest) > 0) { + c.logger.Debugf("Parsed a coalesced packet. Part %d: %d bytes. Remaining: %d bytes.", counter, len(packetData), len(rest)) + } + + p.data = packetData + + processed, err := c.handleLongHeaderPacket(p, hdr, datagramPayloadChecksum) + if err != nil { + return false, err + } + if processed { + wasProcessed = true + } + data = rest + } else { + if counter > 0 { + p.buffer.Split() + } + processed, err := c.handleShortHeaderPacket(p, counter > 0, datagramPayloadChecksum) + if err != nil { + return false, err + } + if processed { + wasProcessed = true + } + break + } + } + + p.buffer.MaybeRelease() + c.blocked = blockModeNone + return wasProcessed, nil +} + +func (c *Conn) handleShortHeaderPacket( + p receivedPacket, + isCoalesced bool, + datagramPayloadChecksum qlog.DatagramPayloadChecksum, // only for logging +) (wasProcessed bool, _ error) { + var wasQueued bool + + defer func() { + // Put back the packet buffer if the packet wasn't queued for later decryption. + if !wasQueued { + p.buffer.Decrement() + } + }() + + destConnID, err := wire.ParseConnectionID(p.data, c.srcConnIDLen) + if err != nil { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType1RTT, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: len(p.data)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + return false, nil + } + pn, pnLen, keyPhase, data, err := c.unpacker.UnpackShortHeader(p.rcvTime, p.data) + if err != nil { + // Stateless reset packets (see RFC 9000, section 10.3): + // * fill the entire UDP datagram (i.e. they cannot be part of a coalesced packet) + // * are short header packets (first bit is 0) + // * have the QUIC bit set (second bit is 1) + // * are at least 21 bytes long + if !isCoalesced && len(p.data) >= protocol.MinReceivedStatelessResetSize && p.data[0]&0b11000000 == 0b01000000 { + token := protocol.StatelessResetToken(p.data[len(p.data)-16:]) + if c.connIDManager.IsActiveStatelessResetToken(token) { + return false, &StatelessResetError{} + } + } + wasQueued, err = c.handleUnpackError(err, p, qlog.PacketType1RTT, datagramPayloadChecksum) + return false, err + } + c.largestRcvdAppData = max(c.largestRcvdAppData, pn) + + if c.logger.Debug() { + c.logger.Debugf("<- Reading packet %d (%d bytes) for connection %s, 1-RTT", pn, p.Size(), destConnID) + wire.LogShortHeader(c.logger, destConnID, pn, pnLen, keyPhase) + } + + if c.receivedPacketHandler.IsPotentiallyDuplicate(pn, protocol.Encryption1RTT) { + c.logger.Debugf("Dropping (potentially) duplicate packet.") + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType1RTT, + PacketNumber: pn, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropDuplicate, + }) + } + return false, nil + } + + var log func([]qlog.Frame) + if c.qlogger != nil { + log = func(frames []qlog.Frame) { + c.qlogger.RecordEvent(qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType1RTT, + DestConnectionID: destConnID, + PacketNumber: pn, + KeyPhaseBit: keyPhase, + }, + Raw: qlog.RawInfo{ + Length: int(p.Size()), + PayloadLength: int(p.Size() - wire.ShortHeaderLen(destConnID, pnLen)), + }, + DatagramPayloadChecksum: datagramPayloadChecksum, + Frames: frames, + ECN: toQlogECN(p.ecn), + }) + } + } + isNonProbing, pathChallenge, err := c.handleUnpackedShortHeaderPacket(destConnID, pn, data, p.ecn, p.rcvTime, log) + if err != nil { + return false, err + } + + // In RFC 9000, only the client can migrate between paths. + if c.perspective == protocol.PerspectiveClient { + return true, nil + } + if addrsEqual(p.remoteAddr, c.RemoteAddr()) { + return true, nil + } + if c.config.DisablePathManager { + // for hysteria2 port hopping, direct change remote address without connection migration logic + c.conn.ChangeRemoteAddr(p.remoteAddr, p.info) + return true, nil + } + + var shouldSwitchPath bool + if c.pathManager == nil { + c.pathManager = newPathManager( + c.connIDManager.GetConnIDForPath, + c.connIDManager.RetireConnIDForPath, + c.logger, + ) + } + destConnID, frames, shouldSwitchPath := c.pathManager.HandlePacket(p.remoteAddr, p.rcvTime, pathChallenge, isNonProbing) + if len(frames) > 0 { + probe, buf, err := c.packer.PackPathProbePacket(destConnID, frames, c.version) + if err != nil { + return true, err + } + c.logger.Debugf("sending path probe packet to %s", p.remoteAddr) + c.logShortHeaderPacketWithDatagramPayloadChecksum(probe, protocol.ECNNon, buf.Len(), false, datagramPayloadChecksum) + c.registerPackedShortHeaderPacket(probe, protocol.ECNNon, p.rcvTime) + c.sendQueue.SendProbe(buf, p.remoteAddr, p.info) + } + // We only switch paths in response to the highest-numbered non-probing packet, + // see section 9.3 of RFC 9000. + if !shouldSwitchPath || pn != c.largestRcvdAppData { + return true, nil + } + c.pathManager.SwitchToPath(p.remoteAddr) + c.sentPacketHandler.MigratedPath(p.rcvTime, protocol.ByteCount(c.config.InitialPacketSize)) + maxPacketSize := protocol.ByteCount(protocol.MaxPacketBufferSize) + if c.peerParams.MaxUDPPayloadSize > 0 && c.peerParams.MaxUDPPayloadSize < maxPacketSize { + maxPacketSize = c.peerParams.MaxUDPPayloadSize + } + c.mtuDiscoverer.Reset( + p.rcvTime, + protocol.ByteCount(c.config.InitialPacketSize), + maxPacketSize, + ) + c.conn.ChangeRemoteAddr(p.remoteAddr, p.info) + return true, nil +} + +func (c *Conn) handleLongHeaderPacket(p receivedPacket, hdr *wire.Header, datagramPayloadChecksum qlog.DatagramPayloadChecksum) (wasProcessed bool, _ error) { + var wasQueued bool + + defer func() { + // Put back the packet buffer if the packet wasn't queued for later decryption. + if !wasQueued { + p.buffer.Decrement() + } + }() + + if hdr.Type == protocol.PacketTypeRetry { + return c.handleRetryPacket(hdr, p.data, p.rcvTime), nil + } + + // The server can change the source connection ID with the first Handshake packet. + // After this, all packets with a different source connection have to be ignored. + if c.receivedFirstPacket && hdr.Type == protocol.PacketTypeInitial && hdr.SrcConnectionID != c.handshakeDestConnID { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropUnknownConnectionID, + }) + } + c.logger.Debugf("Dropping Initial packet (%d bytes) with unexpected source connection ID: %s (expected %s)", p.Size(), hdr.SrcConnectionID, c.handshakeDestConnID) + return false, nil + } + // drop 0-RTT packets, if we are a client + if c.perspective == protocol.PerspectiveClient && hdr.Type == protocol.PacketType0RTT { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false, nil + } + + packet, err := c.unpacker.UnpackLongHeader(hdr, p.data) + if err != nil { + wasQueued, err = c.handleUnpackError(err, p, toQlogPacketType(hdr.Type), datagramPayloadChecksum) + return false, err + } + + if c.logger.Debug() { + c.logger.Debugf("<- Reading packet %d (%d bytes) for connection %s, %s", packet.hdr.PacketNumber, p.Size(), hdr.DestConnectionID, packet.encryptionLevel) + packet.hdr.Log(c.logger) + } + + if pn := packet.hdr.PacketNumber; c.receivedPacketHandler.IsPotentiallyDuplicate(pn, packet.encryptionLevel) { + c.logger.Debugf("Dropping (potentially) duplicate packet.") + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: toQlogPacketType(packet.hdr.Type), + DestConnectionID: hdr.DestConnectionID, + SrcConnectionID: hdr.SrcConnectionID, + PacketNumber: pn, + Version: packet.hdr.Version, + }, + Raw: qlog.RawInfo{Length: int(p.Size()), PayloadLength: int(packet.hdr.Length)}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropDuplicate, + }) + } + return false, nil + } + + if err := c.handleUnpackedLongHeaderPacket(packet, p.ecn, p.rcvTime, datagramPayloadChecksum, p.Size()); err != nil { + return false, err + } + return true, nil +} + +func (c *Conn) handleUnpackError(err error, p receivedPacket, pt qlog.PacketType, datagramPayloadChecksum qlog.DatagramPayloadChecksum) (wasQueued bool, _ error) { + switch err { + case handshake.ErrKeysDropped: + if c.qlogger != nil { + connID, _ := wire.ParseConnectionID(p.data, c.srcConnIDLen) + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: pt, + DestConnectionID: connID, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropKeyUnavailable, + }) + } + c.logger.Debugf("Dropping %s packet (%d bytes) because we already dropped the keys.", pt, p.Size()) + return false, nil + case handshake.ErrKeysNotYetAvailable: + // Sealer for this encryption level not yet available. + // Try again later. + c.tryQueueingUndecryptablePacket(p, pt, datagramPayloadChecksum) + return true, nil + case wire.ErrInvalidReservedBits: + return false, &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: err.Error(), + } + case handshake.ErrDecryptionFailed: + // This might be a packet injected by an attacker. Drop it. + if c.qlogger != nil { + connID, _ := wire.ParseConnectionID(p.data, c.srcConnIDLen) + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: pt, + DestConnectionID: connID, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropPayloadDecryptError, + }) + } + c.logger.Debugf("Dropping %s packet (%d bytes) that could not be unpacked. Error: %s", pt, p.Size(), err) + return false, nil + default: + var headerErr *headerParseError + if errors.As(err, &headerErr) { + // This might be a packet injected by an attacker. Drop it. + if c.qlogger != nil { + connID, _ := wire.ParseConnectionID(p.data, c.srcConnIDLen) + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: pt, + DestConnectionID: connID, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + c.logger.Debugf("Dropping %s packet (%d bytes) for which we couldn't unpack the header. Error: %s", pt, p.Size(), err) + return false, nil + } + // This is an error returned by the AEAD (other than ErrDecryptionFailed). + // For example, a PROTOCOL_VIOLATION due to key updates. + return false, err + } +} + +func (c *Conn) handleRetryPacket(hdr *wire.Header, data []byte, rcvTime monotime.Time) bool /* was this a valid Retry */ { + if c.perspective == protocol.PerspectiveServer { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: hdr.SrcConnectionID, + DestConnectionID: hdr.DestConnectionID, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: len(data)}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + c.logger.Debugf("Ignoring Retry.") + return false + } + if c.receivedFirstPacket { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: hdr.SrcConnectionID, + DestConnectionID: hdr.DestConnectionID, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: len(data)}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + c.logger.Debugf("Ignoring Retry, since we already received a packet.") + return false + } + destConnID := c.connIDManager.Get() + if hdr.SrcConnectionID == destConnID { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: hdr.SrcConnectionID, + DestConnectionID: hdr.DestConnectionID, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: len(data)}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + c.logger.Debugf("Ignoring Retry, since the server didn't change the Source Connection ID.") + return false + } + // If a token is already set, this means that we already received a Retry from the server. + // Ignore this Retry packet. + if c.receivedRetry { + c.logger.Debugf("Ignoring Retry, since a Retry was already received.") + return false + } + + tag := handshake.GetRetryIntegrityTag(data[:len(data)-16], destConnID, hdr.Version) + if !bytes.Equal(data[len(data)-16:], tag[:]) { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: hdr.SrcConnectionID, + DestConnectionID: hdr.DestConnectionID, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: len(data)}, + Trigger: qlog.PacketDropPayloadDecryptError, + }) + } + c.logger.Debugf("Ignoring spoofed Retry. Integrity Tag doesn't match.") + return false + } + + newDestConnID := hdr.SrcConnectionID + c.receivedRetry = true + c.sentPacketHandler.ResetForRetry(rcvTime) + c.handshakeDestConnID = newDestConnID + c.retrySrcConnID = &newDestConnID + c.cryptoStreamHandler.ChangeConnectionID(newDestConnID) + c.packer.SetToken(hdr.Token) + c.connIDManager.ChangeInitialConnID(newDestConnID) + + if c.logger.Debug() { + c.logger.Debugf("<- Received Retry:") + (&wire.ExtendedHeader{Header: *hdr}).Log(c.logger) + c.logger.Debugf("Switching destination connection ID to: %s", hdr.SrcConnectionID) + } + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + DestConnectionID: destConnID, + SrcConnectionID: newDestConnID, + Version: hdr.Version, + Token: &qlog.Token{Raw: hdr.Token}, + }, + Raw: qlog.RawInfo{Length: len(data)}, + }) + } + + c.scheduleSending() + return true +} + +func (c *Conn) handleVersionNegotiationPacket(p receivedPacket) error { + if c.perspective == protocol.PerspectiveServer || // servers never receive version negotiation packets + c.receivedFirstPacket || c.versionNegotiated { // ignore delayed / duplicated version negotiation packets + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return nil + } + + src, dest, supportedVersions, err := wire.ParseVersionNegotiationPacket(p.data) + if err != nil { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + c.logger.Debugf("Error parsing Version Negotiation packet: %s", err) + return nil + } + + if slices.Contains(supportedVersions, c.version) { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedVersion, + }) + } + // The Version Negotiation packet contains the version that we offered. + // This might be a packet sent by an attacker, or it was corrupted. + return nil + } + + c.logger.Infof("Received a Version Negotiation packet. Supported Versions: %s", supportedVersions) + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.VersionNegotiationReceived{ + Header: qlog.PacketHeaderVersionNegotiation{ + DestConnectionID: dest, + SrcConnectionID: src, + }, + SupportedVersions: supportedVersions, + }) + } + newVersion, ok := protocol.ChooseSupportedVersion(c.config.Versions, supportedVersions) + if !ok { + c.destroyImpl(&VersionNegotiationError{ + Ours: c.config.Versions, + Theirs: supportedVersions, + }) + c.logger.Infof("No compatible QUIC version found.") + return nil + } + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.VersionInformation{ + ChosenVersion: newVersion, + ClientVersions: c.config.Versions, + ServerVersions: supportedVersions, + }) + } + + c.logger.Infof("Switching to QUIC version %s.", newVersion) + nextPN, _ := c.sentPacketHandler.PeekPacketNumber(protocol.EncryptionInitial) + return &errCloseForRecreating{ + nextPacketNumber: nextPN, + nextVersion: newVersion, + } +} + +func (c *Conn) handleUnpackedLongHeaderPacket( + packet *unpackedPacket, + ecn protocol.ECN, + rcvTime monotime.Time, + datagramPayloadChecksum qlog.DatagramPayloadChecksum, // only for logging + packetSize protocol.ByteCount, // only for logging +) error { + if !c.receivedFirstPacket { + c.receivedFirstPacket = true + if !c.versionNegotiated && c.qlogger != nil { + var clientVersions, serverVersions []Version + switch c.perspective { + case protocol.PerspectiveClient: + clientVersions = c.config.Versions + case protocol.PerspectiveServer: + serverVersions = c.config.Versions + } + c.qlogger.RecordEvent(qlog.VersionInformation{ + ChosenVersion: c.version, + ClientVersions: clientVersions, + ServerVersions: serverVersions, + }) + } + // The server can change the source connection ID with the first Handshake packet. + if c.perspective == protocol.PerspectiveClient && packet.hdr.SrcConnectionID != c.handshakeDestConnID { + cid := packet.hdr.SrcConnectionID + c.logger.Debugf("Received first packet. Switching destination connection ID to: %s", cid) + c.handshakeDestConnID = cid + c.connIDManager.ChangeInitialConnID(cid) + } + // We create the connection as soon as we receive the first packet from the client. + // We do that before authenticating the packet. + // That means that if the source connection ID was corrupted, + // we might have created a connection with an incorrect source connection ID. + // Once we authenticate the first packet, we need to update it. + if c.perspective == protocol.PerspectiveServer { + if packet.hdr.SrcConnectionID != c.handshakeDestConnID { + c.handshakeDestConnID = packet.hdr.SrcConnectionID + c.connIDManager.ChangeInitialConnID(packet.hdr.SrcConnectionID) + } + if c.qlogger != nil { + var srcAddr, destAddr *net.UDPAddr + if addr, ok := c.conn.LocalAddr().(*net.UDPAddr); ok { + srcAddr = addr + } + if addr, ok := c.conn.RemoteAddr().(*net.UDPAddr); ok { + destAddr = addr + } + c.qlogger.RecordEvent(startedConnectionEvent(srcAddr, destAddr)) + } + } + } + + if c.perspective == protocol.PerspectiveServer && packet.encryptionLevel == protocol.EncryptionHandshake && + !c.droppedInitialKeys { + // On the server side, Initial keys are dropped as soon as the first Handshake packet is received. + // See Section 4.9.1 of RFC 9001. + if err := c.dropEncryptionLevel(protocol.EncryptionInitial, rcvTime); err != nil { + return err + } + } + + c.lastPacketReceivedTime = rcvTime + c.firstAckElicitingPacketAfterIdleSentTime = 0 + c.keepAlivePingSent = false + + if packet.hdr.Type == protocol.PacketType0RTT { + c.largestRcvdAppData = max(c.largestRcvdAppData, packet.hdr.PacketNumber) + } + + var log func([]qlog.Frame) + if c.qlogger != nil { + log = func(frames []qlog.Frame) { + var token *qlog.Token + if len(packet.hdr.Token) > 0 { + token = &qlog.Token{Raw: packet.hdr.Token} + } + c.qlogger.RecordEvent(qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: toQlogPacketType(packet.hdr.Type), + DestConnectionID: packet.hdr.DestConnectionID, + SrcConnectionID: packet.hdr.SrcConnectionID, + PacketNumber: packet.hdr.PacketNumber, + Version: packet.hdr.Version, + Token: token, + }, + Raw: qlog.RawInfo{ + Length: int(packetSize), + PayloadLength: int(packet.hdr.Length), + }, + DatagramPayloadChecksum: datagramPayloadChecksum, + Frames: frames, + ECN: toQlogECN(ecn), + }) + } + } + isAckEliciting, _, _, err := c.handleFrames(packet.data, packet.hdr.DestConnectionID, packet.encryptionLevel, log, rcvTime) + if err != nil { + return err + } + c.sentPacketHandler.ReceivedPacket(packet.encryptionLevel, rcvTime) + return c.receivedPacketHandler.ReceivedPacket(packet.hdr.PacketNumber, ecn, packet.encryptionLevel, rcvTime, isAckEliciting) +} + +func (c *Conn) handleUnpackedShortHeaderPacket( + destConnID protocol.ConnectionID, + pn protocol.PacketNumber, + data []byte, + ecn protocol.ECN, + rcvTime monotime.Time, + log func([]qlog.Frame), +) (isNonProbing bool, pathChallenge *wire.PathChallengeFrame, _ error) { + c.lastPacketReceivedTime = rcvTime + c.firstAckElicitingPacketAfterIdleSentTime = 0 + c.keepAlivePingSent = false + + isAckEliciting, isNonProbing, pathChallenge, err := c.handleFrames(data, destConnID, protocol.Encryption1RTT, log, rcvTime) + if err != nil { + return false, nil, err + } + c.sentPacketHandler.ReceivedPacket(protocol.Encryption1RTT, rcvTime) + if err := c.receivedPacketHandler.ReceivedPacket(pn, ecn, protocol.Encryption1RTT, rcvTime, isAckEliciting); err != nil { + return false, nil, err + } + return isNonProbing, pathChallenge, nil +} + +// handleFrames parses the frames, one after the other, and handles them. +// It returns the last PATH_CHALLENGE frame contained in the packet, if any. +func (c *Conn) handleFrames( + data []byte, + destConnID protocol.ConnectionID, + encLevel protocol.EncryptionLevel, + log func([]qlog.Frame), + rcvTime monotime.Time, +) (isAckEliciting, isNonProbing bool, pathChallenge *wire.PathChallengeFrame, _ error) { + // Only used for tracing. + // If we're not tracing, this slice will always remain empty. + var frames []qlog.Frame + if log != nil { + frames = make([]qlog.Frame, 0, 4) + } + handshakeWasComplete := c.handshakeComplete + var handleErr error + var skipHandling bool + + for len(data) > 0 { + frameType, l, err := c.frameParser.ParseType(data, encLevel) + if err != nil { + // The frame parser skips over PADDING frames, and returns an io.EOF if the PADDING + // frames were the last frames in this packet. + if err == io.EOF { + break + } + return false, false, nil, err + } + data = data[l:] + + if ackhandler.IsFrameTypeAckEliciting(frameType) { + isAckEliciting = true + } + if !wire.IsProbingFrameType(frameType) { + isNonProbing = true + } + + // We're inlining common cases, to avoid using interfaces + // Fast path: STREAM, DATAGRAM and ACK + if frameType.IsStreamFrameType() { + streamFrame, l, err := c.frameParser.ParseStreamFrame(frameType, data, c.version) + if err != nil { + return false, false, nil, err + } + data = data[l:] + + if log != nil { + frames = append(frames, toQlogFrame(streamFrame)) + } + // an error occurred handling a previous frame, don't handle the current frame + if skipHandling { + continue + } + wire.LogFrame(c.logger, streamFrame, false) + handleErr = c.streamsMap.HandleStreamFrame(streamFrame, rcvTime) + } else if frameType.IsAckFrameType() { + ackFrame, l, err := c.frameParser.ParseAckFrame(frameType, data, encLevel, c.version) + if err != nil { + return false, false, nil, err + } + data = data[l:] + if log != nil { + frames = append(frames, toQlogFrame(ackFrame)) + } + // an error occurred handling a previous frame, don't handle the current frame + if skipHandling { + continue + } + wire.LogFrame(c.logger, ackFrame, false) + handleErr = c.handleAckFrame(ackFrame, encLevel, rcvTime) + } else if frameType.IsDatagramFrameType() { + datagramFrame, l, err := c.frameParser.ParseDatagramFrame(frameType, data, c.version) + if err != nil { + return false, false, nil, err + } + data = data[l:] + + if log != nil { + frames = append(frames, toQlogFrame(datagramFrame)) + } + // an error occurred handling a previous frame, don't handle the current frame + if skipHandling { + continue + } + wire.LogFrame(c.logger, datagramFrame, false) + handleErr = c.handleDatagramFrame(datagramFrame) + } else { + frame, l, err := c.frameParser.ParseLessCommonFrame(frameType, data, c.version) + if err != nil { + return false, false, nil, err + } + data = data[l:] + + if log != nil { + frames = append(frames, toQlogFrame(frame)) + } + // an error occurred handling a previous frame, don't handle the current frame + if skipHandling { + continue + } + pc, err := c.handleFrame(frame, encLevel, destConnID, rcvTime) + if pc != nil { + pathChallenge = pc + } + handleErr = err + } + + if handleErr != nil { + // if we're logging, we need to keep parsing (but not handling) all frames + skipHandling = true + if log == nil { + return false, false, nil, handleErr + } + } + } + + if log != nil { + log(frames) + if handleErr != nil { + return false, false, nil, handleErr + } + } + + // Handle completion of the handshake after processing all the frames. + // This ensures that we correctly handle the following case on the server side: + // We receive a Handshake packet that contains the CRYPTO frame that allows us to complete the handshake, + // and an ACK serialized after that CRYPTO frame. In this case, we still want to process the ACK frame. + if !handshakeWasComplete && c.handshakeComplete { + if err := c.handleHandshakeComplete(rcvTime); err != nil { + return false, false, nil, err + } + } + return +} + +func (c *Conn) handleFrame( + f wire.Frame, + encLevel protocol.EncryptionLevel, + destConnID protocol.ConnectionID, + rcvTime monotime.Time, +) (pathChallenge *wire.PathChallengeFrame, _ error) { + var err error + wire.LogFrame(c.logger, f, false) + switch frame := f.(type) { + case *wire.CryptoFrame: + err = c.handleCryptoFrame(frame, encLevel, rcvTime) + case *wire.ConnectionCloseFrame: + err = c.handleConnectionCloseFrame(frame) + case *wire.ResetStreamFrame: + err = c.streamsMap.HandleResetStreamFrame(frame, rcvTime) + case *wire.MaxDataFrame: + c.connFlowController.UpdateSendWindow(frame.MaximumData) + case *wire.MaxStreamDataFrame: + err = c.streamsMap.HandleMaxStreamDataFrame(frame) + case *wire.MaxStreamsFrame: + c.streamsMap.HandleMaxStreamsFrame(frame) + case *wire.DataBlockedFrame: + case *wire.StreamDataBlockedFrame: + err = c.streamsMap.HandleStreamDataBlockedFrame(frame) + case *wire.StreamsBlockedFrame: + case *wire.StopSendingFrame: + err = c.streamsMap.HandleStopSendingFrame(frame) + case *wire.PingFrame: + case *wire.PathChallengeFrame: + c.handlePathChallengeFrame(frame) + pathChallenge = frame + case *wire.PathResponseFrame: + err = c.handlePathResponseFrame(frame) + case *wire.NewTokenFrame: + err = c.handleNewTokenFrame(frame) + case *wire.NewConnectionIDFrame: + err = c.connIDManager.Add(frame) + case *wire.RetireConnectionIDFrame: + err = c.connIDGenerator.Retire(frame.SequenceNumber, destConnID, rcvTime.Add(3*c.rttStats.PTO(false))) + case *wire.HandshakeDoneFrame: + err = c.handleHandshakeDoneFrame(rcvTime) + default: + err = fmt.Errorf("unexpected frame type: %s", reflect.ValueOf(&frame).Elem().Type().Name()) + } + return pathChallenge, err +} + +// handlePacket is called by the server with a new packet +func (c *Conn) handlePacket(p receivedPacket) { + c.receivedPacketMx.Lock() + // Discard packets once the amount of queued packets is larger than + // the channel size, protocol.MaxConnUnprocessedPackets + if c.receivedPackets.Len() >= protocol.MaxConnUnprocessedPackets { + if c.qlogger != nil { + var datagramPayloadChecksum qlog.DatagramPayloadChecksum + if wire.IsLongHeaderPacket(p.data[0]) { + datagramPayloadChecksum = qlog.CalculateDatagramPayloadChecksum(p.data) + } + c.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + c.receivedPacketMx.Unlock() + return + } + c.receivedPackets.PushBack(p) + c.receivedPacketMx.Unlock() + + select { + case c.notifyReceivedPacket <- struct{}{}: + default: + } +} + +func (c *Conn) handleConnectionCloseFrame(frame *wire.ConnectionCloseFrame) error { + if frame.IsApplicationError { + return &qerr.ApplicationError{ + Remote: true, + ErrorCode: qerr.ApplicationErrorCode(frame.ErrorCode), + ErrorMessage: frame.ReasonPhrase, + } + } + return &qerr.TransportError{ + Remote: true, + ErrorCode: qerr.TransportErrorCode(frame.ErrorCode), + FrameType: frame.FrameType, + ErrorMessage: frame.ReasonPhrase, + } +} + +func (c *Conn) handleCryptoFrame(frame *wire.CryptoFrame, encLevel protocol.EncryptionLevel, rcvTime monotime.Time) error { + if err := c.cryptoStreamManager.HandleCryptoFrame(frame, encLevel); err != nil { + return err + } + for { + data := c.cryptoStreamManager.GetCryptoData(encLevel) + if data == nil { + break + } + if err := c.cryptoStreamHandler.HandleMessage(data, encLevel); err != nil { + return err + } + } + return c.handleHandshakeEvents(rcvTime) +} + +func (c *Conn) handleHandshakeEvents(now monotime.Time) error { + for { + ev := c.cryptoStreamHandler.NextEvent() + var err error + switch ev.Kind { + case handshake.EventNoEvent: + return nil + case handshake.EventHandshakeComplete: + // Don't call handleHandshakeComplete yet. + // It's advantageous to process ACK frames that might be serialized after the CRYPTO frame first. + c.handshakeComplete = true + case handshake.EventReceivedTransportParameters: + err = c.handleTransportParameters(ev.TransportParameters) + case handshake.EventRestoredTransportParameters: + c.restoreTransportParameters(ev.TransportParameters) + close(c.earlyConnReadyChan) + case handshake.EventReceivedReadKeys: + // New keys mean the encryption level we send at is about to change, at + // which point the imitated client sizes the packet number against the + // full datagram again rather than the room left by the last one. + c.sentPacketHandler.SetLastDatagramPadding(0) + // queue all previously undecryptable packets + c.undecryptablePacketsToProcess = append(c.undecryptablePacketsToProcess, c.undecryptablePackets...) + c.undecryptablePackets = nil + case handshake.EventDiscard0RTTKeys: + err = c.dropEncryptionLevel(protocol.Encryption0RTT, now) + case handshake.EventWriteInitialData: + _, err = c.initialStream.Write(ev.Data) + case handshake.EventWriteHandshakeData: + _, err = c.handshakeStream.Write(ev.Data) + } + if err != nil { + return err + } + } +} + +func (c *Conn) handlePathChallengeFrame(f *wire.PathChallengeFrame) { + if c.perspective == protocol.PerspectiveClient { + c.queueControlFrame(&wire.PathResponseFrame{Data: f.Data}) + } +} + +func (c *Conn) handlePathResponseFrame(f *wire.PathResponseFrame) error { + switch c.perspective { + case protocol.PerspectiveClient: + return c.handlePathResponseFrameClient(f) + case protocol.PerspectiveServer: + return c.handlePathResponseFrameServer(f) + default: + panic("unreachable") + } +} + +func (c *Conn) handlePathResponseFrameClient(f *wire.PathResponseFrame) error { + pm := c.pathManagerOutgoing.Load() + if pm == nil { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "unexpected PATH_RESPONSE frame", + } + } + pm.HandlePathResponseFrame(f) + return nil +} + +func (c *Conn) handlePathResponseFrameServer(f *wire.PathResponseFrame) error { + if c.pathManager == nil { + // since we didn't send PATH_CHALLENGEs yet, we don't expect PATH_RESPONSEs + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "unexpected PATH_RESPONSE frame", + } + } + c.pathManager.HandlePathResponseFrame(f) + return nil +} + +func (c *Conn) handleNewTokenFrame(frame *wire.NewTokenFrame) error { + if c.perspective == protocol.PerspectiveServer { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "received NEW_TOKEN frame from the client", + } + } + if c.config.TokenStore != nil { + c.config.TokenStore.Put(c.tokenStoreKey, &ClientToken{data: frame.Token, rtt: c.rttStats.SmoothedRTT()}) + } + return nil +} + +func (c *Conn) handleHandshakeDoneFrame(rcvTime monotime.Time) error { + if c.perspective == protocol.PerspectiveServer { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "received a HANDSHAKE_DONE frame", + } + } + if !c.handshakeConfirmed { + return c.handleHandshakeConfirmed(rcvTime) + } + return nil +} + +func (c *Conn) handleAckFrame(frame *wire.AckFrame, encLevel protocol.EncryptionLevel, rcvTime monotime.Time) error { + acked1RTTPacket, err := c.sentPacketHandler.ReceivedAck(frame, encLevel, c.lastPacketReceivedTime) + if err != nil { + return err + } + if !acked1RTTPacket { + return nil + } + // On the client side: If the packet acknowledged a 1-RTT packet, this confirms the handshake. + // This is only possible if the ACK was sent in a 1-RTT packet. + // This is an optimization over simply waiting for a HANDSHAKE_DONE frame, see section 4.1.2 of RFC 9001. + if c.perspective == protocol.PerspectiveClient && !c.handshakeConfirmed { + if err := c.handleHandshakeConfirmed(rcvTime); err != nil { + return err + } + } + // If one of the acknowledged packets was a Path MTU probe packet, this might have increased the Path MTU estimate. + if c.mtuDiscoverer != nil { + mtu := c.mtuDiscoverer.CurrentSize() + maxPayloadSize := estimateMaxPayloadSize(mtu) + if maxPayloadSize > protocol.ByteCount(c.maxPayloadSizeEstimate.Load()) { + c.maxPayloadSizeEstimate.Store(uint32(maxPayloadSize)) + c.sentPacketHandler.SetMaxDatagramSize(mtu) + } + } + return c.cryptoStreamHandler.SetLargest1RTTAcked(frame.LargestAcked()) +} + +func (c *Conn) handleDatagramFrame(f *wire.DatagramFrame) error { + if f.Length(c.version) > wire.MaxDatagramSize { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "DATAGRAM frame too large", + } + } + c.datagramQueue.HandleDatagramFrame(f) + return nil +} + +func (c *Conn) setCloseError(e *closeError) { + c.closeErr.CompareAndSwap(nil, e) + select { + case c.closeChan <- struct{}{}: + default: + } +} + +// closeLocal closes the connection and send a CONNECTION_CLOSE containing the error +func (c *Conn) closeLocal(e error) { + c.setCloseError(&closeError{err: e, immediate: false}) +} + +// destroy closes the connection without sending the error on the wire +func (c *Conn) destroy(e error) { + c.destroyImpl(e) + <-c.ctx.Done() +} + +func (c *Conn) destroyImpl(e error) { + c.setCloseError(&closeError{err: e, immediate: true}) +} + +// CloseWithError closes the connection with an error. +// The error string will be sent to the peer. +func (c *Conn) CloseWithError(code ApplicationErrorCode, desc string) error { + c.closeLocal(&qerr.ApplicationError{ + ErrorCode: code, + ErrorMessage: desc, + }) + <-c.ctx.Done() + return nil +} + +func (c *Conn) closeWithTransportError(code TransportErrorCode) { + c.closeLocal(&qerr.TransportError{ErrorCode: code}) + <-c.ctx.Done() +} + +func (c *Conn) handleCloseError(closeErr *closeError) { + if closeErr.immediate { + if nerr, ok := closeErr.err.(net.Error); ok && nerr.Timeout() { + c.logger.Errorf("Destroying connection: %s", closeErr.err) + } else { + c.logger.Errorf("Destroying connection with error: %s", closeErr.err) + } + } else { + if closeErr.err == nil { + c.logger.Infof("Closing connection.") + } else { + c.logger.Errorf("Closing connection with error: %s", closeErr.err) + } + } + + e := closeErr.err + if e == nil { + e = &qerr.ApplicationError{} + } else { + defer func() { closeErr.err = e }() + } + + var ( + statelessResetErr *StatelessResetError + versionNegotiationErr *VersionNegotiationError + recreateErr *errCloseForRecreating + applicationErr *ApplicationError + transportErr *TransportError + ) + var isRemoteClose bool + var trigger qlog.ConnectionCloseTrigger + var reason string + var transportErrorCode *qlog.TransportErrorCode + var applicationErrorCode *qlog.ApplicationErrorCode + switch { + case errors.Is(e, qerr.ErrIdleTimeout), + errors.Is(e, qerr.ErrHandshakeTimeout): + trigger = qlog.ConnectionCloseTriggerIdleTimeout + case errors.As(e, &statelessResetErr): + trigger = qlog.ConnectionCloseTriggerStatelessReset + case errors.As(e, &versionNegotiationErr): + trigger = qlog.ConnectionCloseTriggerVersionMismatch + case errors.As(e, &recreateErr): + case errors.As(e, &applicationErr): + isRemoteClose = applicationErr.Remote + reason = applicationErr.ErrorMessage + applicationErrorCode = &applicationErr.ErrorCode + case errors.As(e, &transportErr): + isRemoteClose = transportErr.Remote + reason = transportErr.ErrorMessage + transportErrorCode = &transportErr.ErrorCode + case closeErr.immediate: + e = closeErr.err + default: + te := &qerr.TransportError{ + ErrorCode: qerr.InternalError, + ErrorMessage: e.Error(), + } + e = te + reason = te.ErrorMessage + code := te.ErrorCode + transportErrorCode = &code + } + + c.streamsMap.CloseWithError(e) + if c.datagramQueue != nil { + c.datagramQueue.CloseWithError(e) + } + + // In rare instances, the connection ID manager might switch to a new connection ID + // when sending the CONNECTION_CLOSE frame. + // The connection ID manager removes the active stateless reset token from the packet + // handler map when it is closed, so we need to make sure that this happens last. + defer c.connIDManager.Close() + + if c.qlogger != nil && !errors.As(e, &recreateErr) { + initiator := qlog.InitiatorLocal + if isRemoteClose { + initiator = qlog.InitiatorRemote + } + c.qlogger.RecordEvent(qlog.ConnectionClosed{ + Initiator: initiator, + ConnectionError: transportErrorCode, + ApplicationError: applicationErrorCode, + Trigger: trigger, + Reason: reason, + }) + } + + // If this is a remote close we're done here + if isRemoteClose { + c.connIDGenerator.ReplaceWithClosed(nil, 3*c.rttStats.PTO(false)) + return + } + if closeErr.immediate { + c.connIDGenerator.RemoveAll() + return + } + // Don't send out any CONNECTION_CLOSE if this is an error that occurred + // before we even sent out the first packet. + if c.perspective == protocol.PerspectiveClient && !c.sentFirstPacket { + c.connIDGenerator.RemoveAll() + return + } + connClosePacket, err := c.sendConnectionClose(e) + if err != nil { + c.logger.Debugf("Error sending CONNECTION_CLOSE: %s", err) + } + c.connIDGenerator.ReplaceWithClosed(connClosePacket, 3*c.rttStats.PTO(false)) +} + +func (c *Conn) dropEncryptionLevel(encLevel protocol.EncryptionLevel, now monotime.Time) error { + c.sentPacketHandler.DropPackets(encLevel, now) + c.receivedPacketHandler.DropPackets(encLevel) + //nolint:exhaustive // only Initial and 0-RTT need special treatment + switch encLevel { + case protocol.EncryptionInitial: + c.droppedInitialKeys = true + c.cryptoStreamHandler.DiscardInitialKeys() + case protocol.Encryption0RTT: + c.streamsMap.ResetFor0RTT() + c.framer.Handle0RTTRejection() + return c.connFlowController.Reset() + } + return c.cryptoStreamManager.Drop(encLevel) +} + +// is called for the client, when restoring transport parameters saved for 0-RTT +func (c *Conn) restoreTransportParameters(params *wire.TransportParameters) { + if c.logger.Debug() { + c.logger.Debugf("Restoring Transport Parameters: %s", params) + } + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.ParametersSet{ + Restore: true, + Initiator: qlog.InitiatorRemote, + SentBy: c.perspective, + OriginalDestinationConnectionID: params.OriginalDestinationConnectionID, + InitialSourceConnectionID: params.InitialSourceConnectionID, + RetrySourceConnectionID: params.RetrySourceConnectionID, + StatelessResetToken: params.StatelessResetToken, + DisableActiveMigration: params.DisableActiveMigration, + MaxIdleTimeout: params.MaxIdleTimeout, + MaxUDPPayloadSize: params.MaxUDPPayloadSize, + AckDelayExponent: params.AckDelayExponent, + MaxAckDelay: params.MaxAckDelay, + ActiveConnectionIDLimit: params.ActiveConnectionIDLimit, + InitialMaxData: params.InitialMaxData, + InitialMaxStreamDataBidiLocal: params.InitialMaxStreamDataBidiLocal, + InitialMaxStreamDataBidiRemote: params.InitialMaxStreamDataBidiRemote, + InitialMaxStreamDataUni: params.InitialMaxStreamDataUni, + InitialMaxStreamsBidi: int64(params.MaxBidiStreamNum), + InitialMaxStreamsUni: int64(params.MaxUniStreamNum), + MaxDatagramFrameSize: params.MaxDatagramFrameSize, + EnableResetStreamAt: params.EnableResetStreamAt, + }) + } + + c.peerParams = params + c.connIDGenerator.SetMaxActiveConnIDs(params.ActiveConnectionIDLimit) + c.connFlowController.UpdateSendWindow(params.InitialMaxData) + c.streamsMap.HandleTransportParameters(params) +} + +func (c *Conn) handleTransportParameters(params *wire.TransportParameters) error { + if c.qlogger != nil { + c.qlogTransportParameters(params, c.perspective.Opposite(), false) + } + if err := c.checkTransportParameters(params); err != nil { + return &qerr.TransportError{ + ErrorCode: qerr.TransportParameterError, + ErrorMessage: err.Error(), + } + } + + if c.perspective == protocol.PerspectiveClient && c.peerParams != nil && c.ConnectionState().Used0RTT && !params.ValidForUpdate(c.peerParams) { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "server sent reduced limits after accepting 0-RTT data", + } + } + + c.peerParams = params + // On the client side we have to wait for handshake completion. + // During a 0-RTT connection, we are only allowed to use the new transport parameters for 1-RTT packets. + if c.perspective == protocol.PerspectiveServer { + c.applyTransportParameters() + // On the server side, the early connection is ready as soon as we processed + // the client's transport parameters. + close(c.earlyConnReadyChan) + } + return nil +} + +func (c *Conn) checkTransportParameters(params *wire.TransportParameters) error { + if c.logger.Debug() { + c.logger.Debugf("Processed Transport Parameters: %s", params) + } + + // check the initial_source_connection_id + if params.InitialSourceConnectionID != c.handshakeDestConnID { + return fmt.Errorf("expected initial_source_connection_id to equal %s, is %s", c.handshakeDestConnID, params.InitialSourceConnectionID) + } + + if c.perspective == protocol.PerspectiveServer { + return nil + } + // check the original_destination_connection_id + if params.OriginalDestinationConnectionID != c.origDestConnID { + return fmt.Errorf("expected original_destination_connection_id to equal %s, is %s", c.origDestConnID, params.OriginalDestinationConnectionID) + } + if c.retrySrcConnID != nil { // a Retry was performed + if params.RetrySourceConnectionID == nil { + return errors.New("missing retry_source_connection_id") + } + if *params.RetrySourceConnectionID != *c.retrySrcConnID { + return fmt.Errorf("expected retry_source_connection_id to equal %s, is %s", c.retrySrcConnID, *params.RetrySourceConnectionID) + } + } else if params.RetrySourceConnectionID != nil { + return errors.New("received retry_source_connection_id, although no Retry was performed") + } + return nil +} + +func (c *Conn) applyTransportParameters() { + params := c.peerParams + // Our local idle timeout will always be > 0. + c.idleTimeout = c.config.MaxIdleTimeout + // If the peer advertised an idle timeout, take the minimum of the values. + if params.MaxIdleTimeout > 0 { + c.idleTimeout = min(c.idleTimeout, params.MaxIdleTimeout) + } + c.keepAliveInterval = min(c.config.KeepAlivePeriod, c.idleTimeout/2) + c.streamsMap.HandleTransportParameters(params) + c.frameParser.SetAckDelayExponent(params.AckDelayExponent) + c.connFlowController.UpdateSendWindow(params.InitialMaxData) + c.rttStats.SetMaxAckDelay(params.MaxAckDelay) + c.connIDGenerator.SetMaxActiveConnIDs(params.ActiveConnectionIDLimit) + if params.StatelessResetToken != nil { + c.connIDManager.SetStatelessResetToken(*params.StatelessResetToken) + } + // We don't support connection migration yet, so we don't have any use for the preferred_address. + if params.PreferredAddress != nil { + // Retire the connection ID. + c.connIDManager.AddFromPreferredAddress(params.PreferredAddress.ConnectionID, params.PreferredAddress.StatelessResetToken) + } + maxPacketSize := protocol.ByteCount(protocol.MaxPacketBufferSize) + if params.MaxUDPPayloadSize > 0 && params.MaxUDPPayloadSize < maxPacketSize { + maxPacketSize = params.MaxUDPPayloadSize + } + c.mtuDiscoverer = newMTUDiscoverer( + c.rttStats, + protocol.ByteCount(c.config.InitialPacketSize), + maxPacketSize, + c.qlogger, + ) +} + +func (c *Conn) triggerSending(now monotime.Time) error { + c.pacingDeadline = 0 + + sendMode := c.sentPacketHandler.SendMode(now) + switch sendMode { + case ackhandler.SendAny: + return c.sendPackets(now) + case ackhandler.SendNone: + c.blocked = blockModeHardBlocked + return nil + case ackhandler.SendPacingLimited: + deadline := c.sentPacketHandler.TimeUntilSend() + if deadline.IsZero() { + deadline = deadlineSendImmediately + } + c.pacingDeadline = deadline + // Allow sending of an ACK if we're pacing limit. + // This makes sure that a peer that is mostly receiving data (and thus has an inaccurate cwnd estimate) + // sends enough ACKs to allow its peer to utilize the bandwidth. + return c.maybeSendAckOnlyPacket(now) + case ackhandler.SendAck: + // We can at most send a single ACK only packet. + // There will only be a new ACK after receiving new packets. + // SendAck is only returned when we're congestion limited, so we don't need to set the pacing timer. + c.blocked = blockModeCongestionLimited + return c.maybeSendAckOnlyPacket(now) + case ackhandler.SendPTOInitial, ackhandler.SendPTOHandshake, ackhandler.SendPTOAppData: + if err := c.sendProbePacket(sendMode, now); err != nil { + return err + } + if c.sendQueue.WouldBlock() { + c.scheduleSending() + return nil + } + return c.triggerSending(now) + default: + return fmt.Errorf("BUG: invalid send mode %d", sendMode) + } +} + +func (c *Conn) sendPackets(now monotime.Time) error { + if c.perspective == protocol.PerspectiveClient && c.handshakeConfirmed { + if pm := c.pathManagerOutgoing.Load(); pm != nil { + connID, frame, tr, ok := pm.NextPathToProbe() + if ok { + probe, buf, err := c.packer.PackPathProbePacket(connID, []ackhandler.Frame{frame}, c.version) + if err != nil { + return err + } + c.logger.Debugf("sending path probe packet from %s", c.LocalAddr()) + c.logShortHeaderPacket(probe, protocol.ECNNon, buf.Len()) + c.registerPackedShortHeaderPacket(probe, protocol.ECNNon, now) + tr.WriteTo(buf.Data, c.conn.RemoteAddr()) + // There's (likely) more data to send. Loop around again. + c.scheduleSending() + return nil + } + } + } + + // Path MTU Discovery + // Can't use GSO, since we need to send a single packet that's larger than our current maximum size. + // Performance-wise, this doesn't matter, since we only send a very small (<10) number of + // MTU probe packets per connection. + if c.handshakeConfirmed && c.mtuDiscoverer != nil && c.mtuDiscoverer.ShouldSendProbe(now) { + ping, size := c.mtuDiscoverer.GetPing(now) + p, buf, err := c.packer.PackMTUProbePacket(ping, size, c.version) + if err != nil { + return err + } + ecn := c.sentPacketHandler.ECNMode(true) + c.logShortHeaderPacket(p, ecn, buf.Len()) + c.registerPackedShortHeaderPacket(p, ecn, now) + c.sendQueue.Send(buf, 0, ecn) + // There's (likely) more data to send. Loop around again. + c.scheduleSending() + return nil + } + + if offset := c.connFlowController.GetWindowUpdate(now); offset > 0 { + c.framer.QueueControlFrame(&wire.MaxDataFrame{MaximumData: offset}) + } + if cf := c.cryptoStreamManager.GetPostHandshakeData(protocol.MaxPostHandshakeCryptoFrameSize); cf != nil { + c.queueControlFrame(cf) + } + + if !c.handshakeConfirmed { + packet, err := c.packer.PackCoalescedPacket(false, c.maxPacketSize(), now, c.version) + if err != nil || packet == nil { + return err + } + c.sentFirstPacket = true + if err := c.sendPackedCoalescedPacket(packet, c.sentPacketHandler.ECNMode(packet.IsOnlyShortHeaderPacket()), now); err != nil { + return err + } + //nolint:exhaustive // only need to handle pacing-related events here + switch c.sentPacketHandler.SendMode(now) { + case ackhandler.SendPacingLimited: + c.resetPacingDeadline() + case ackhandler.SendAny: + c.pacingDeadline = deadlineSendImmediately + } + return nil + } + + if c.conn.capabilities().GSO { + return c.sendPacketsWithGSO(now) + } + return c.sendPacketsWithoutGSO(now) +} + +func (c *Conn) sendPacketsWithoutGSO(now monotime.Time) error { + for { + buf := getPacketBuffer() + ecn := c.sentPacketHandler.ECNMode(true) + if _, err := c.appendOneShortHeaderPacket(buf, c.maxPacketSize(), ecn, now); err != nil { + if err == errNothingToPack { + buf.Release() + return nil + } + return err + } + + c.sendQueue.Send(buf, 0, ecn) + + if c.sendQueue.WouldBlock() { + return nil + } + sendMode := c.sentPacketHandler.SendMode(now) + if sendMode == ackhandler.SendPacingLimited { + c.resetPacingDeadline() + return nil + } + if sendMode != ackhandler.SendAny { + return nil + } + // Prioritize receiving of packets over sending out more packets. + c.receivedPacketMx.Lock() + hasPackets := !c.receivedPackets.Empty() + c.receivedPacketMx.Unlock() + if hasPackets { + c.pacingDeadline = deadlineSendImmediately + return nil + } + } +} + +func (c *Conn) sendPacketsWithGSO(now monotime.Time) error { + buf := getLargePacketBuffer() + maxSize := c.maxPacketSize() + + ecn := c.sentPacketHandler.ECNMode(true) + for { + var dontSendMore bool + size, err := c.appendOneShortHeaderPacket(buf, maxSize, ecn, now) + if err != nil { + if err != errNothingToPack { + return err + } + if buf.Len() == 0 { + buf.Release() + return nil + } + dontSendMore = true + } + + if !dontSendMore { + sendMode := c.sentPacketHandler.SendMode(now) + if sendMode == ackhandler.SendPacingLimited { + c.resetPacingDeadline() + } + if sendMode != ackhandler.SendAny { + dontSendMore = true + } + } + + // Don't send more packets in this batch if they require a different ECN marking than the previous ones. + nextECN := c.sentPacketHandler.ECNMode(true) + + // Append another packet if + // 1. The congestion controller and pacer allow sending more + // 2. The last packet appended was a full-size packet + // 3. The next packet will have the same ECN marking + // 4. We still have enough space for another full-size packet in the buffer + if !dontSendMore && size == maxSize && nextECN == ecn && buf.Len()+maxSize <= buf.Cap() { + continue + } + + c.sendQueue.Send(buf, uint16(maxSize), ecn) + + if dontSendMore { + return nil + } + if c.sendQueue.WouldBlock() { + return nil + } + + // Prioritize receiving of packets over sending out more packets. + c.receivedPacketMx.Lock() + hasPackets := !c.receivedPackets.Empty() + c.receivedPacketMx.Unlock() + if hasPackets { + c.pacingDeadline = deadlineSendImmediately + return nil + } + + ecn = nextECN + buf = getLargePacketBuffer() + } +} + +func (c *Conn) resetPacingDeadline() { + deadline := c.sentPacketHandler.TimeUntilSend() + if deadline.IsZero() { + deadline = deadlineSendImmediately + } + c.pacingDeadline = deadline +} + +func (c *Conn) maybeSendAckOnlyPacket(now monotime.Time) error { + if !c.handshakeConfirmed { + ecn := c.sentPacketHandler.ECNMode(false) + packet, err := c.packer.PackCoalescedPacket(true, c.maxPacketSize(), now, c.version) + if err != nil { + return err + } + if packet == nil { + return nil + } + return c.sendPackedCoalescedPacket(packet, ecn, now) + } + + ecn := c.sentPacketHandler.ECNMode(true) + p, buf, err := c.packer.PackAckOnlyPacket(c.maxPacketSize(), now, c.version) + if err != nil { + if err == errNothingToPack { + return nil + } + return err + } + c.logShortHeaderPacket(p, ecn, buf.Len()) + c.registerPackedShortHeaderPacket(p, ecn, now) + c.sendQueue.Send(buf, 0, ecn) + return nil +} + +func (c *Conn) sendProbePacket(sendMode ackhandler.SendMode, now monotime.Time) error { + var encLevel protocol.EncryptionLevel + //nolint:exhaustive // We only need to handle the PTO send modes here. + switch sendMode { + case ackhandler.SendPTOInitial: + encLevel = protocol.EncryptionInitial + case ackhandler.SendPTOHandshake: + encLevel = protocol.EncryptionHandshake + case ackhandler.SendPTOAppData: + encLevel = protocol.Encryption1RTT + default: + return fmt.Errorf("connection BUG: unexpected send mode: %d", sendMode) + } + // Queue probe packets until we actually send out a packet, + // or until there are no more packets to queue. + var packet *coalescedPacket + for packet == nil { + if wasQueued := c.sentPacketHandler.QueueProbePacket(encLevel); !wasQueued { + break + } + var err error + packet, err = c.packer.PackPTOProbePacket(encLevel, c.maxPacketSize(), false, now, c.version) + if err != nil { + return err + } + } + if packet == nil { + var err error + packet, err = c.packer.PackPTOProbePacket(encLevel, c.maxPacketSize(), true, now, c.version) + if err != nil { + return err + } + } + if packet == nil || (len(packet.longHdrPackets) == 0 && packet.shortHdrPacket == nil) { + return fmt.Errorf("connection BUG: couldn't pack %s probe packet: %v", encLevel, packet) + } + return c.sendPackedCoalescedPacket(packet, c.sentPacketHandler.ECNMode(packet.IsOnlyShortHeaderPacket()), now) +} + +// appendOneShortHeaderPacket appends a new packet to the given packetBuffer. +// If there was nothing to pack, the returned size is 0. +func (c *Conn) appendOneShortHeaderPacket(buf *packetBuffer, maxSize protocol.ByteCount, ecn protocol.ECN, now monotime.Time) (protocol.ByteCount, error) { + startLen := buf.Len() + p, err := c.packer.AppendPacket(buf, maxSize, now, c.version) + if err != nil { + return 0, err + } + size := buf.Len() - startLen + c.logShortHeaderPacket(p, ecn, size) + c.registerPackedShortHeaderPacket(p, ecn, now) + return size, nil +} + +func (c *Conn) registerPackedShortHeaderPacket(p shortHeaderPacket, ecn protocol.ECN, now monotime.Time) { + if p.IsPathProbePacket { + c.sentPacketHandler.SentPacket( + now, + p.PacketNumber, + protocol.InvalidPacketNumber, + p.StreamFrames, + p.Frames, + protocol.Encryption1RTT, + ecn, + p.Length, + p.IsPathMTUProbePacket, + true, + ) + return + } + if c.firstAckElicitingPacketAfterIdleSentTime.IsZero() && (len(p.StreamFrames) > 0 || ackhandler.HasAckElicitingFrames(p.Frames)) { + c.firstAckElicitingPacketAfterIdleSentTime = now + } + + largestAcked := protocol.InvalidPacketNumber + if p.Ack != nil { + largestAcked = p.Ack.LargestAcked() + } + c.sentPacketHandler.SentPacket( + now, + p.PacketNumber, + largestAcked, + p.StreamFrames, + p.Frames, + protocol.Encryption1RTT, + ecn, + p.Length, + p.IsPathMTUProbePacket, + false, + ) + c.connIDManager.SentPacket() +} + +func (c *Conn) sendPackedCoalescedPacket(packet *coalescedPacket, ecn protocol.ECN, now monotime.Time) error { + c.logCoalescedPacket(packet, ecn) + for _, p := range packet.longHdrPackets { + if c.firstAckElicitingPacketAfterIdleSentTime.IsZero() && p.IsAckEliciting() { + c.firstAckElicitingPacketAfterIdleSentTime = now + } + largestAcked := protocol.InvalidPacketNumber + if p.ack != nil { + largestAcked = p.ack.LargestAcked() + } + c.sentPacketHandler.SentPacket( + now, + p.header.PacketNumber, + largestAcked, + p.streamFrames, + p.frames, + p.EncryptionLevel(), + ecn, + p.length, + false, + false, + ) + if c.perspective == protocol.PerspectiveClient && p.EncryptionLevel() == protocol.EncryptionHandshake && + !c.droppedInitialKeys { + // On the client side, Initial keys are dropped as soon as the first Handshake packet is sent. + // See Section 4.9.1 of RFC 9001. + if err := c.dropEncryptionLevel(protocol.EncryptionInitial, now); err != nil { + return err + } + } + } + if p := packet.shortHdrPacket; p != nil { + if c.firstAckElicitingPacketAfterIdleSentTime.IsZero() && p.IsAckEliciting() { + c.firstAckElicitingPacketAfterIdleSentTime = now + } + largestAcked := protocol.InvalidPacketNumber + if p.Ack != nil { + largestAcked = p.Ack.LargestAcked() + } + c.sentPacketHandler.SentPacket( + now, + p.PacketNumber, + largestAcked, + p.StreamFrames, + p.Frames, + protocol.Encryption1RTT, + ecn, + p.Length, + p.IsPathMTUProbePacket, + false, + ) + } + c.connIDManager.SentPacket() + c.sendQueue.Send(packet.buffer, 0, ecn) + return nil +} + +func (c *Conn) sendConnectionClose(e error) ([]byte, error) { + var packet *coalescedPacket + var err error + var transportErr *qerr.TransportError + var applicationErr *qerr.ApplicationError + if errors.As(e, &transportErr) { + packet, err = c.packer.PackConnectionClose(transportErr, c.maxPacketSize(), c.version) + } else if errors.As(e, &applicationErr) { + packet, err = c.packer.PackApplicationClose(applicationErr, c.maxPacketSize(), c.version) + } else { + packet, err = c.packer.PackConnectionClose(&qerr.TransportError{ + ErrorCode: qerr.InternalError, + ErrorMessage: fmt.Sprintf("connection BUG: unspecified error type (msg: %s)", e.Error()), + }, c.maxPacketSize(), c.version) + } + if err != nil { + return nil, err + } + ecn := c.sentPacketHandler.ECNMode(packet.IsOnlyShortHeaderPacket()) + c.logCoalescedPacket(packet, ecn) + return packet.buffer.Data, c.conn.Write(packet.buffer.Data, 0, ecn) +} + +func (c *Conn) maxPacketSize() protocol.ByteCount { + if c.mtuDiscoverer == nil { + // Use the configured packet size on the client side. + // If the server sends a max_udp_payload_size that's smaller than this size, we can ignore this: + // Apparently the server still processed the (fully padded) Initial packet anyway. + if c.perspective == protocol.PerspectiveClient { + return protocol.ByteCount(c.config.InitialPacketSize) + } + // On the server side, there's no downside to using 1200 bytes until we received the client's transport + // parameters: + // * If the first packet didn't contain the entire ClientHello, all we can do is ACK that packet. We don't + // need a lot of bytes for that. + // * If it did, we will have processed the transport parameters and initialized the MTU discoverer. + return protocol.MinInitialPacketSize + } + return c.mtuDiscoverer.CurrentSize() +} + +// AcceptStream returns the next stream opened by the peer, blocking until one is available. +func (c *Conn) AcceptStream(ctx context.Context) (*Stream, error) { + return c.streamsMap.AcceptStream(ctx) +} + +// AcceptUniStream returns the next unidirectional stream opened by the peer, blocking until one is available. +func (c *Conn) AcceptUniStream(ctx context.Context) (*ReceiveStream, error) { + return c.streamsMap.AcceptUniStream(ctx) +} + +// OpenStream opens a new bidirectional QUIC stream. +// There is no signaling to the peer about new streams: +// The peer can only accept the stream after data has been sent on the stream, +// or the stream has been reset or closed. +// When reaching the peer's stream limit, it is not possible to open a new stream until the +// peer raises the stream limit. In that case, a [StreamLimitReachedError] is returned. +func (c *Conn) OpenStream() (*Stream, error) { + return c.streamsMap.OpenStream() +} + +// OpenStreamSync opens a new bidirectional QUIC stream. +// It blocks until a new stream can be opened. +// There is no signaling to the peer about new streams: +// The peer can only accept the stream after data has been sent on the stream, +// or the stream has been reset or closed. +func (c *Conn) OpenStreamSync(ctx context.Context) (*Stream, error) { + return c.streamsMap.OpenStreamSync(ctx) +} + +// OpenUniStream opens a new outgoing unidirectional QUIC stream. +// There is no signaling to the peer about new streams: +// The peer can only accept the stream after data has been sent on the stream, +// or the stream has been reset or closed. +// When reaching the peer's stream limit, it is not possible to open a new stream until the +// peer raises the stream limit. In that case, a [StreamLimitReachedError] is returned. +func (c *Conn) OpenUniStream() (*SendStream, error) { + return c.streamsMap.OpenUniStream() +} + +// OpenUniStreamSync opens a new outgoing unidirectional QUIC stream. +// It blocks until a new stream can be opened. +// There is no signaling to the peer about new streams: +// The peer can only accept the stream after data has been sent on the stream, +// or the stream has been reset or closed. +func (c *Conn) OpenUniStreamSync(ctx context.Context) (*SendStream, error) { + return c.streamsMap.OpenUniStreamSync(ctx) +} + +func (c *Conn) newFlowController(id protocol.StreamID) *streamFlowController { + initialSendWindow := c.peerParams.InitialMaxStreamDataUni + if protocol.StreamTypeOf(id) == protocol.StreamTypeBidi { + if protocol.StreamInitiator(id) == c.perspective { + initialSendWindow = c.peerParams.InitialMaxStreamDataBidiRemote + } else { + initialSendWindow = c.peerParams.InitialMaxStreamDataBidiLocal + } + } + return newStreamFlowController( + id, + c.connFlowController, + protocol.ByteCount(c.config.InitialStreamReceiveWindow), + protocol.ByteCount(c.config.MaxStreamReceiveWindow), + initialSendWindow, + c.rttStats, + c.logger, + ) +} + +// scheduleSending signals that we have data for sending +func (c *Conn) scheduleSending() { + select { + case c.sendingScheduled <- struct{}{}: + default: + } +} + +// tryQueueingUndecryptablePacket queues a packet for which we're missing the decryption keys. +// The qlogevents.PacketType is only used for logging purposes. +func (c *Conn) tryQueueingUndecryptablePacket(p receivedPacket, pt qlog.PacketType, datagramPayloadChecksum qlog.DatagramPayloadChecksum) { + if c.handshakeComplete { + panic("shouldn't queue undecryptable packets after handshake completion") + } + if len(c.undecryptablePackets)+1 > protocol.MaxUndecryptablePackets { + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: pt, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + c.logger.Infof("Dropping undecryptable packet (%d bytes). Undecryptable packet queue full.", p.Size()) + return + } + c.logger.Infof("Queueing packet (%d bytes) for later decryption", p.Size()) + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.PacketBuffered{ + Header: qlog.PacketHeader{ + PacketType: pt, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum, + }) + } + c.undecryptablePackets = append(c.undecryptablePackets, receivedPacketWithChecksum{receivedPacket: p, checksum: datagramPayloadChecksum}) +} + +func (c *Conn) queueControlFrame(f wire.Frame) { + c.framer.QueueControlFrame(f) + c.scheduleSending() +} + +func (c *Conn) onHasConnectionData() { c.scheduleSending() } + +func (c *Conn) onHasStreamData(id protocol.StreamID, str *SendStream) { + c.framer.AddActiveStream(id, str) + c.scheduleSending() +} + +func (c *Conn) onHasStreamControlFrame(id protocol.StreamID, str streamControlFrameGetter) { + c.framer.AddStreamWithControlFrames(id, str) + c.scheduleSending() +} + +func (c *Conn) onStreamCompleted(id protocol.StreamID) { + if err := c.streamsMap.DeleteStream(id); err != nil { + c.closeLocal(err) + } + c.framer.RemoveActiveStream(id) +} + +// SendDatagram sends a message using a QUIC datagram, as specified in RFC 9221, +// if the peer enabled datagram support. +// There is no delivery guarantee for DATAGRAM frames, they are not retransmitted if lost. +// The payload of the datagram needs to fit into a single QUIC packet. +// In addition, a datagram may be dropped before being sent out if the available packet size suddenly decreases. +// If the payload is too large to be sent at the current time, a DatagramTooLargeError is returned. +func (c *Conn) SendDatagram(p []byte) error { + if !c.supportsDatagrams() { + return errors.New("datagram support disabled") + } + + f := &wire.DatagramFrame{DataLenPresent: true} + maxDatagramFrameSize := c.peerMaxDatagramFrameSize() + // The payload size estimate is conservative. + // Under many circumstances we could send a few more bytes. + maxDataLen := min( + f.MaxDataLen(maxDatagramFrameSize, c.version), + protocol.ByteCount(c.maxPayloadSizeEstimate.Load()), + ) + if protocol.ByteCount(len(p)) > maxDataLen { + return &DatagramTooLargeError{MaxDatagramPayloadSize: int64(maxDataLen)} + } + f.Data = make([]byte, len(p)) + copy(f.Data, p) + return c.datagramQueue.Add(f) +} + +// ReceiveDatagram gets a message received in a QUIC datagram, as specified in RFC 9221. +func (c *Conn) ReceiveDatagram(ctx context.Context) ([]byte, error) { + if !c.config.EnableDatagrams { + return nil, errors.New("datagram support disabled") + } + return c.datagramQueue.Receive(ctx) +} + +// LocalAddr returns the local address of the QUIC connection. +func (c *Conn) LocalAddr() net.Addr { return c.conn.LocalAddr() } + +// RemoteAddr returns the remote address of the QUIC connection. +func (c *Conn) RemoteAddr() net.Addr { return c.conn.RemoteAddr() } + +// getPathManager lazily initializes the Conn's pathManagerOutgoing. +// May create multiple pathManagerOutgoing objects if called concurrently. +func (c *Conn) getPathManager() *pathManagerOutgoing { + old := c.pathManagerOutgoing.Load() + if old != nil { + // Path manager is already initialized + return old + } + + // Initialize the path manager + new := newPathManagerOutgoing( + c.connIDManager.GetConnIDForPath, + c.connIDManager.RetireConnIDForPath, + c.scheduleSending, + ) + if c.pathManagerOutgoing.CompareAndSwap(old, new) { + return new + } + + // Swap failed. A concurrent writer wrote first, use their value. + return c.pathManagerOutgoing.Load() +} + +func (c *Conn) AddPath(t *Transport) (*Path, error) { + if c.perspective == protocol.PerspectiveServer { + return nil, errors.New("server cannot initiate connection migration") + } + if c.peerParams.DisableActiveMigration { + return nil, errors.New("server disabled connection migration") + } + if err := t.init(false); err != nil { + return nil, err + } + return c.getPathManager().NewPath( + t, + 200*time.Millisecond, // initial RTT estimate + func() { + runner := (*packetHandlerMap)(t) + c.connIDGenerator.AddConnRunner( + runner, + connRunnerCallbacks{ + AddConnectionID: func(connID protocol.ConnectionID) { runner.Add(connID, c) }, + RemoveConnectionID: runner.Remove, + ReplaceWithClosed: runner.ReplaceWithClosed, + }, + ) + }, + ), nil +} + +// HandshakeComplete blocks until the handshake completes (or fails). +// For the client, data sent before completion of the handshake is encrypted with 0-RTT keys. +// For the server, data sent before completion of the handshake is encrypted with 1-RTT keys, +// however the client's identity is only verified once the handshake completes. +func (c *Conn) HandshakeComplete() <-chan struct{} { + return c.handshakeCompleteChan +} + +// QlogTrace returns the qlog trace of the QUIC connection. +// It is nil if qlog is not enabled. +func (c *Conn) QlogTrace() qlogwriter.Trace { + return c.qlogTrace +} + +// NextConnection transitions a connection to be usable after a 0-RTT rejection. +// It waits for the handshake to complete and then enables the connection for normal use. +// This should be called when the server rejects 0-RTT and the application receives +// [Err0RTTRejected] errors. +// +// Note that 0-RTT rejection invalidates all data sent in 0-RTT packets. It is the +// application's responsibility to handle this (for example by resending the data). +func (c *Conn) NextConnection(ctx context.Context) (*Conn, error) { + // The handshake might fail after the server rejected 0-RTT. + // This could happen if the Finished message is malformed or never received. + select { + case <-ctx.Done(): + return nil, context.Cause(ctx) + case <-c.Context().Done(): + case <-c.HandshakeComplete(): + c.streamsMap.UseResetMaps() + } + return c, nil +} + +// estimateMaxPayloadSize estimates the maximum payload size for short header packets. +// It is not very sophisticated: it just subtracts the size of header (assuming the maximum +// connection ID length), and the size of the encryption tag. +func estimateMaxPayloadSize(mtu protocol.ByteCount) protocol.ByteCount { + return mtu - 1 /* type byte */ - 20 /* maximum connection ID length */ - 16 /* tag size */ +} + +// SetCongestionControl replace the current congestion control algorithm with a new one. +func (c *Conn) SetCongestionControl(cc congestion.CongestionControl) { + c.sentPacketHandler.SetCongestionControl(cc) +} + +// InitialPacketSize returns the datagram size the connection starts out with, +// before path MTU discovery raises it. This is the value the connection's own +// congestion controller is seeded with, and it is not always the package +// default: options that make the connection start smaller lower it too. +// +// A controller installed with SetCongestionControl must be seeded with this, +// not with the package default. Seeding it high instead breaks as soon as path +// MTU discovery reports a size between the two, which the connection sees as an +// increase but the controller sees as a decrease. +func (c *Conn) InitialPacketSize() congestion.ByteCount { + return congestion.ByteCount(c.config.InitialPacketSize) +} + +// SetRemoteAddr Replace the current remote addr with a new one +func (c *Conn) SetRemoteAddr(addr net.Addr) { + c.conn.SetRemoteAddr(addr) +} + +// Config Return current config +func (c *Conn) Config() *Config { + return c.config +} diff --git a/third_party/quic-go/connection_logging.go b/third_party/quic-go/connection_logging.go new file mode 100644 index 0000000..c465ace --- /dev/null +++ b/third_party/quic-go/connection_logging.go @@ -0,0 +1,315 @@ +package quic + +import ( + "net" + "net/netip" + "slices" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" +) + +// ConvertFrame converts a wire.Frame into a logging.Frame. +// This makes it possible for external packages to access the frames. +// Furthermore, it removes the data slices from CRYPTO and STREAM frames. +func toQlogFrame(frame wire.Frame) qlog.Frame { + switch f := frame.(type) { + case *wire.AckFrame: + // We use a pool for ACK frames. + // Implementations of the tracer interface may hold on to frames, so we need to make a copy here. + return qlog.Frame{Frame: toQlogAckFrame(f)} + case *wire.CryptoFrame: + return qlog.Frame{ + Frame: &qlog.CryptoFrame{ + Offset: int64(f.Offset), + Length: int64(len(f.Data)), + }, + } + case *wire.StreamFrame: + return qlog.Frame{ + Frame: &qlog.StreamFrame{ + StreamID: f.StreamID, + Offset: int64(f.Offset), + Length: int64(f.DataLen()), + Fin: f.Fin, + }, + } + case *wire.DatagramFrame: + return qlog.Frame{ + Frame: &qlog.DatagramFrame{ + Length: int64(len(f.Data)), + }, + } + default: + return qlog.Frame{Frame: frame} + } +} + +func toQlogAckFrame(f *wire.AckFrame) *qlog.AckFrame { + ack := &qlog.AckFrame{ + AckRanges: slices.Clone(f.AckRanges), + DelayTime: f.DelayTime, + ECNCE: f.ECNCE, + ECT0: f.ECT0, + ECT1: f.ECT1, + } + return ack +} + +func (c *Conn) logLongHeaderPacket(p *longHeaderPacket, ecn protocol.ECN, datagramPayloadChecksum qlog.DatagramPayloadChecksum) { + // quic-go logging + if c.logger.Debug() { + p.header.Log(c.logger) + if p.ack != nil { + wire.LogFrame(c.logger, p.ack, true) + } + for _, frame := range p.frames { + wire.LogFrame(c.logger, frame.Frame, true) + } + for _, frame := range p.streamFrames { + wire.LogFrame(c.logger, frame.Frame, true) + } + } + + // tracing + if c.qlogger != nil { + numFrames := len(p.frames) + len(p.streamFrames) + if p.ack != nil { + numFrames++ + } + frames := make([]qlog.Frame, 0, numFrames) + if p.ack != nil { + frames = append(frames, toQlogFrame(p.ack)) + } + for _, f := range p.frames { + frames = append(frames, toQlogFrame(f.Frame)) + } + for _, f := range p.streamFrames { + frames = append(frames, toQlogFrame(f.Frame)) + } + c.qlogger.RecordEvent(qlog.PacketSent{ + Header: qlog.PacketHeader{ + PacketType: toQlogPacketType(p.header.Type), + KeyPhaseBit: p.header.KeyPhase, + PacketNumber: p.header.PacketNumber, + Version: p.header.Version, + SrcConnectionID: p.header.SrcConnectionID, + DestConnectionID: p.header.DestConnectionID, + }, + Raw: qlog.RawInfo{ + Length: int(p.length), + PayloadLength: int(p.header.Length), + }, + DatagramPayloadChecksum: datagramPayloadChecksum, + Frames: frames, + ECN: toQlogECN(ecn), + }) + } +} + +func (c *Conn) logShortHeaderPacket(p shortHeaderPacket, ecn protocol.ECN, size protocol.ByteCount) { + c.logShortHeaderPacketWithDatagramPayloadChecksum(p, ecn, size, false, 0) +} + +func (c *Conn) logShortHeaderPacketWithDatagramPayloadChecksum(p shortHeaderPacket, ecn protocol.ECN, size protocol.ByteCount, isCoalesced bool, datagramPayloadChecksum qlog.DatagramPayloadChecksum) { + if c.logger.Debug() && !isCoalesced { + c.logger.Debugf("-> Sending packet %d (%d bytes) for connection %s, 1-RTT (ECN: %s)", p.PacketNumber, size, c.logID, ecn) + } + // quic-go logging + if c.logger.Debug() { + wire.LogShortHeader(c.logger, p.DestConnID, p.PacketNumber, p.PacketNumberLen, p.KeyPhase) + if p.Ack != nil { + wire.LogFrame(c.logger, p.Ack, true) + } + for _, f := range p.Frames { + wire.LogFrame(c.logger, f.Frame, true) + } + for _, f := range p.StreamFrames { + wire.LogFrame(c.logger, f.Frame, true) + } + } + + // tracing + if c.qlogger != nil { + numFrames := len(p.Frames) + len(p.StreamFrames) + if p.Ack != nil { + numFrames++ + } + fs := make([]qlog.Frame, 0, numFrames) + if p.Ack != nil { + fs = append(fs, toQlogFrame(p.Ack)) + } + for _, f := range p.Frames { + fs = append(fs, toQlogFrame(f.Frame)) + } + for _, f := range p.StreamFrames { + fs = append(fs, toQlogFrame(f.Frame)) + } + c.qlogger.RecordEvent(qlog.PacketSent{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType1RTT, + KeyPhaseBit: p.KeyPhase, + PacketNumber: p.PacketNumber, + Version: c.version, + DestConnectionID: p.DestConnID, + }, + Raw: qlog.RawInfo{ + Length: int(size), + PayloadLength: int(size - wire.ShortHeaderLen(p.DestConnID, p.PacketNumberLen)), + }, + DatagramPayloadChecksum: datagramPayloadChecksum, + Frames: fs, + ECN: toQlogECN(ecn), + }) + } +} + +func (c *Conn) logCoalescedPacket(packet *coalescedPacket, ecn protocol.ECN) { + var datagramPayloadChecksum qlog.DatagramPayloadChecksum + if c.qlogger != nil { + datagramPayloadChecksum = qlog.CalculateDatagramPayloadChecksum(packet.buffer.Data) + } + if c.logger.Debug() { + // There's a short period between dropping both Initial and Handshake keys and completion of the handshake, + // during which we might call PackCoalescedPacket but just pack a short header packet. + if len(packet.longHdrPackets) == 0 && packet.shortHdrPacket != nil { + c.logShortHeaderPacketWithDatagramPayloadChecksum( + *packet.shortHdrPacket, + ecn, + packet.shortHdrPacket.Length, + false, + datagramPayloadChecksum, + ) + return + } + if len(packet.longHdrPackets) > 1 { + c.logger.Debugf("-> Sending coalesced packet (%d parts, %d bytes) for connection %s", len(packet.longHdrPackets), packet.buffer.Len(), c.logID) + } else { + c.logger.Debugf("-> Sending packet %d (%d bytes) for connection %s, %s", packet.longHdrPackets[0].header.PacketNumber, packet.buffer.Len(), c.logID, packet.longHdrPackets[0].EncryptionLevel()) + } + } + for _, p := range packet.longHdrPackets { + c.logLongHeaderPacket(p, ecn, datagramPayloadChecksum) + } + if p := packet.shortHdrPacket; p != nil { + c.logShortHeaderPacketWithDatagramPayloadChecksum(*p, ecn, p.Length, true, datagramPayloadChecksum) + } +} + +func (c *Conn) qlogTransportParameters(tp *wire.TransportParameters, sentBy protocol.Perspective, restore bool) { + ev := qlog.ParametersSet{ + Restore: restore, + OriginalDestinationConnectionID: tp.OriginalDestinationConnectionID, + InitialSourceConnectionID: tp.InitialSourceConnectionID, + RetrySourceConnectionID: tp.RetrySourceConnectionID, + StatelessResetToken: tp.StatelessResetToken, + DisableActiveMigration: tp.DisableActiveMigration, + MaxIdleTimeout: tp.MaxIdleTimeout, + MaxUDPPayloadSize: tp.MaxUDPPayloadSize, + AckDelayExponent: tp.AckDelayExponent, + MaxAckDelay: tp.MaxAckDelay, + ActiveConnectionIDLimit: tp.ActiveConnectionIDLimit, + InitialMaxData: tp.InitialMaxData, + InitialMaxStreamDataBidiLocal: tp.InitialMaxStreamDataBidiLocal, + InitialMaxStreamDataBidiRemote: tp.InitialMaxStreamDataBidiRemote, + InitialMaxStreamDataUni: tp.InitialMaxStreamDataUni, + InitialMaxStreamsBidi: int64(tp.MaxBidiStreamNum), + InitialMaxStreamsUni: int64(tp.MaxUniStreamNum), + MaxDatagramFrameSize: tp.MaxDatagramFrameSize, + EnableResetStreamAt: tp.EnableResetStreamAt, + } + if sentBy == c.perspective { + ev.Initiator = qlog.InitiatorLocal + } else { + ev.Initiator = qlog.InitiatorRemote + } + if tp.PreferredAddress != nil { + ev.PreferredAddress = &qlog.PreferredAddress{ + IPv4: tp.PreferredAddress.IPv4, + IPv6: tp.PreferredAddress.IPv6, + ConnectionID: tp.PreferredAddress.ConnectionID, + StatelessResetToken: tp.PreferredAddress.StatelessResetToken, + } + } + c.qlogger.RecordEvent(ev) +} + +func toQlogECN(ecn protocol.ECN) qlog.ECN { + //nolint:exhaustive // only need to handle the 3 valid values + switch ecn { + case protocol.ECT0: + return qlog.ECT0 + case protocol.ECT1: + return qlog.ECT1 + case protocol.ECNCE: + return qlog.ECNCE + default: + return qlog.ECNUnsupported + } +} + +func toQlogPacketType(pt protocol.PacketType) qlog.PacketType { + var qpt qlog.PacketType + switch pt { + case protocol.PacketTypeInitial: + qpt = qlog.PacketTypeInitial + case protocol.PacketTypeHandshake: + qpt = qlog.PacketTypeHandshake + case protocol.PacketType0RTT: + qpt = qlog.PacketType0RTT + case protocol.PacketTypeRetry: + qpt = qlog.PacketTypeRetry + } + return qpt +} + +func toPathEndpointInfo(addr *net.UDPAddr) qlog.PathEndpointInfo { + if addr == nil { + return qlog.PathEndpointInfo{} + } + + var info qlog.PathEndpointInfo + if addr.IP == nil || addr.IP.To4() != nil { + addrPort := netip.AddrPortFrom(netip.AddrFrom4([4]byte(addr.IP.To4())), uint16(addr.Port)) + if addrPort.IsValid() { + info.IPv4 = addrPort + } + } else { + addrPort := netip.AddrPortFrom(netip.AddrFrom16([16]byte(addr.IP.To16())), uint16(addr.Port)) + if addrPort.IsValid() { + info.IPv6 = addrPort + } + } + return info +} + +// startedConnectionEvent builds a StartedConnection event using consistent logic +// for both endpoints. If the local address is unspecified (e.g., dual-stack +// listener), it selects the family based on the remote address and uses the +// unspecified address of that family with the local port. +func startedConnectionEvent(local, remote *net.UDPAddr) qlog.StartedConnection { + var localInfo, remoteInfo qlog.PathEndpointInfo + if remote != nil { + remoteInfo = toPathEndpointInfo(remote) + } + if local != nil { + if local.IP == nil || local.IP.IsUnspecified() { + // Choose local family based on the remote address family. + if remote != nil && remote.IP.To4() != nil { + ap := netip.AddrPortFrom(netip.AddrFrom4([4]byte{}), uint16(local.Port)) + if ap.IsValid() { + localInfo.IPv4 = ap + } + } else if remote != nil && remote.IP.To16() != nil && remote.IP.To4() == nil { + ap := netip.AddrPortFrom(netip.AddrFrom16([16]byte{}), uint16(local.Port)) + if ap.IsValid() { + localInfo.IPv6 = ap + } + } + } else { + localInfo = toPathEndpointInfo(local) + } + } + return qlog.StartedConnection{Local: localInfo, Remote: remoteInfo} +} diff --git a/third_party/quic-go/connection_logging_test.go b/third_party/quic-go/connection_logging_test.go new file mode 100644 index 0000000..7ce04cc --- /dev/null +++ b/third_party/quic-go/connection_logging_test.go @@ -0,0 +1,143 @@ +package quic + +import ( + "net" + "net/netip" + "testing" + + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + + "github.com/stretchr/testify/require" +) + +func TestConnectionLoggingCryptoFrame(t *testing.T) { + f := toQlogFrame(&wire.CryptoFrame{ + Offset: 1234, + Data: []byte("foobar"), + }) + require.Equal(t, &qlog.CryptoFrame{ + Offset: 1234, + Length: 6, + }, f.Frame) +} + +func TestConnectionLoggingStreamFrame(t *testing.T) { + f := toQlogFrame(&wire.StreamFrame{ + StreamID: 42, + Offset: 1234, + Data: []byte("foo"), + Fin: true, + }) + require.Equal(t, &qlog.StreamFrame{ + StreamID: 42, + Offset: 1234, + Length: 3, + Fin: true, + }, f.Frame) +} + +func TestConnectionLoggingAckFrame(t *testing.T) { + ack := &wire.AckFrame{ + AckRanges: []wire.AckRange{ + {Smallest: 1, Largest: 3}, + {Smallest: 6, Largest: 7}, + }, + DelayTime: 42, + ECNCE: 123, + ECT0: 456, + ECT1: 789, + } + f := toQlogFrame(ack) + // now modify the ACK range in the original frame + ack.AckRanges[0].Smallest = 2 + require.Equal(t, &qlog.AckFrame{ + AckRanges: []wire.AckRange{ + {Smallest: 1, Largest: 3}, // unchanged, since the ACK ranges were cloned + {Smallest: 6, Largest: 7}, + }, + DelayTime: 42, + ECNCE: 123, + ECT0: 456, + ECT1: 789, + }, f.Frame) +} + +func TestConnectionLoggingDatagramFrame(t *testing.T) { + f := toQlogFrame(&wire.DatagramFrame{Data: []byte("foobar")}) + require.Equal(t, &qlog.DatagramFrame{Length: 6}, f.Frame) +} + +func TestConnectionLoggingOtherFrames(t *testing.T) { + f := toQlogFrame(&wire.MaxDataFrame{MaximumData: 1234}) + require.Equal(t, &qlog.MaxDataFrame{MaximumData: 1234}, f.Frame) +} + +func TestConnectionLoggingStartedConnectionEvent(t *testing.T) { + tests := []struct { + name string + local *net.UDPAddr + remote *net.UDPAddr + wantLocalIP string + wantLocalPort uint16 + wantRemote netip.AddrPort + }{ + { + name: "unspecified local, remote IPv4 -> 0.0.0.0", + local: &net.UDPAddr{Port: 58451}, + remote: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 6121}, + wantLocalIP: "0.0.0.0", + wantLocalPort: 58451, + wantRemote: netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 6121), + }, + { + name: "unspecified local, remote IPv6 -> ::", + local: &net.UDPAddr{Port: 4242}, + remote: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 6121}, + wantLocalIP: "::", + wantLocalPort: 4242, + wantRemote: func() netip.AddrPort { a, _ := netip.ParseAddr("2001:db8::1"); return netip.AddrPortFrom(a, 6121) }(), + }, + { + name: "specified local IPv4", + local: &net.UDPAddr{IP: net.IPv4(192, 168, 1, 10), Port: 9999}, + remote: &net.UDPAddr{IP: net.IPv4(10, 0, 0, 1), Port: 1234}, + wantLocalIP: "192.168.1.10", + wantLocalPort: 9999, + wantRemote: netip.AddrPortFrom(netip.AddrFrom4([4]byte{10, 0, 0, 1}), 1234), + }, + { + name: "specified local IPv6", + local: &net.UDPAddr{IP: net.ParseIP("fe80::1"), Port: 999}, + remote: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 6121}, + wantLocalIP: "fe80::1", + wantLocalPort: 999, + wantRemote: func() netip.AddrPort { a, _ := netip.ParseAddr("2001:db8::1"); return netip.AddrPortFrom(a, 6121) }(), + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ev := startedConnectionEvent(tc.local, tc.remote) + var gotIP string + var gotPort uint16 + if ev.Local.IPv4.IsValid() { + gotIP = ev.Local.IPv4.Addr().String() + gotPort = ev.Local.IPv4.Port() + } else if ev.Local.IPv6.IsValid() { + gotIP = ev.Local.IPv6.Addr().String() + gotPort = ev.Local.IPv6.Port() + } + require.Equal(t, tc.wantLocalIP, gotIP) + require.Equal(t, tc.wantLocalPort, gotPort) + + var gotRemote netip.AddrPort + if ev.Remote.IPv4.IsValid() { + gotRemote = ev.Remote.IPv4 + } else if ev.Remote.IPv6.IsValid() { + gotRemote = ev.Remote.IPv6 + } + require.Equal(t, tc.wantRemote, gotRemote) + }) + } +} diff --git a/third_party/quic-go/connection_test.go b/third_party/quic-go/connection_test.go new file mode 100644 index 0000000..d5ead48 --- /dev/null +++ b/third_party/quic-go/connection_test.go @@ -0,0 +1,3485 @@ +package quic + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/tls" + "errors" + "fmt" + "net" + "net/netip" + "strconv" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/mocks" + mockackhandler "github.com/apernet/quic-go/internal/mocks/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type testConnectionOpt func(*Conn) + +func connectionOptCryptoSetup(cs *mocks.MockCryptoSetup) testConnectionOpt { + return func(conn *Conn) { conn.cryptoStreamHandler = cs } +} + +func connectionOptConnFlowController(cfc *connectionFlowController) testConnectionOpt { + return func(conn *Conn) { conn.connFlowController = cfc } +} + +func connectionOptTracer(r qlogwriter.Recorder) testConnectionOpt { + return func(conn *Conn) { conn.qlogger = r } +} + +func connectionOptSentPacketHandler(sph ackhandler.SentPacketHandler) testConnectionOpt { + return func(conn *Conn) { conn.sentPacketHandler = sph } +} + +func connectionOptUnpacker(u unpacker) testConnectionOpt { + return func(conn *Conn) { conn.unpacker = u } +} + +func connectionOptSender(s sender) testConnectionOpt { + return func(conn *Conn) { conn.sendQueue = s } +} + +func connectionOptHandshakeConfirmed() testConnectionOpt { + return func(conn *Conn) { + conn.handshakeComplete = true + conn.handshakeConfirmed = true + } +} + +func connectionOptRTT(rtt time.Duration) testConnectionOpt { + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + return func(conn *Conn) { conn.rttStats = rttStats } +} + +func connectionOptRetrySrcConnID(rcid protocol.ConnectionID) testConnectionOpt { + return func(conn *Conn) { conn.retrySrcConnID = &rcid } +} + +type testConnection struct { + conn *Conn + connRunner *MockConnRunner + sendConn *MockSendConn + packer *MockPacker + destConnID protocol.ConnectionID + srcConnID protocol.ConnectionID + remoteAddr *net.UDPAddr +} + +func (tc *testConnection) receivedPacketHandler() *ackhandler.ReceivedPacketHandler { + return &tc.conn.receivedPacketHandler +} + +func newServerTestConnection( + t *testing.T, + mockCtrl *gomock.Controller, + config *Config, + gso bool, + opts ...testConnectionOpt, +) *testConnection { + if mockCtrl == nil { + mockCtrl = gomock.NewController(t) + } + remoteAddr := &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 4321} + localAddr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 1234} + connRunner := NewMockConnRunner(mockCtrl) + sendConn := NewMockSendConn(mockCtrl) + sendConn.EXPECT().capabilities().Return(connCapabilities{GSO: gso}).AnyTimes() + sendConn.EXPECT().RemoteAddr().Return(remoteAddr).AnyTimes() + sendConn.EXPECT().LocalAddr().Return(localAddr).AnyTimes() + packer := NewMockPacker(mockCtrl) + b := make([]byte, 12) + rand.Read(b) + origDestConnID := protocol.ParseConnectionID(b[:6]) + srcConnID := protocol.ParseConnectionID(b[6:12]) + ctx, cancel := context.WithCancelCause(context.Background()) + if config == nil { + config = &Config{DisablePathMTUDiscovery: true} + } + wc := newConnection( + ctx, + cancel, + sendConn, + connRunner, + origDestConnID, + nil, + protocol.ConnectionID{}, + protocol.ConnectionID{}, + srcConnID, + &protocol.DefaultConnectionIDGenerator{}, + newStatelessResetter(nil), + populateConfig(config), + &tls.Config{}, + handshake.NewTokenGenerator(handshake.TokenProtectorKey{}), + false, + 1337*time.Millisecond, + nil, + utils.DefaultLogger, + protocol.Version1, + ) + require.Nil(t, wc.testHooks) + conn := wc.Conn + conn.packer = packer + for _, opt := range opts { + opt(conn) + } + return &testConnection{ + conn: conn, + connRunner: connRunner, + sendConn: sendConn, + packer: packer, + destConnID: origDestConnID, + srcConnID: srcConnID, + remoteAddr: remoteAddr, + } +} + +func newClientTestConnection( + t *testing.T, + mockCtrl *gomock.Controller, + config *Config, + enable0RTT bool, + opts ...testConnectionOpt, +) *testConnection { + if mockCtrl == nil { + mockCtrl = gomock.NewController(t) + } + remoteAddr := &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 4321} + localAddr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 1234} + connRunner := NewMockConnRunner(mockCtrl) + sendConn := NewMockSendConn(mockCtrl) + sendConn.EXPECT().capabilities().Return(connCapabilities{}).AnyTimes() + sendConn.EXPECT().RemoteAddr().Return(remoteAddr).AnyTimes() + sendConn.EXPECT().LocalAddr().Return(localAddr).AnyTimes() + packer := NewMockPacker(mockCtrl) + b := make([]byte, 12) + rand.Read(b) + destConnID := protocol.ParseConnectionID(b[:6]) + srcConnID := protocol.ParseConnectionID(b[6:12]) + if config == nil { + config = &Config{DisablePathMTUDiscovery: true} + } + conn, err := newClientConnection( + context.Background(), + sendConn, + connRunner, + destConnID, + srcConnID, + &protocol.DefaultConnectionIDGenerator{}, + newStatelessResetter(nil), + populateConfig(config), + &tls.Config{ServerName: "quic-go.net"}, + 0, + enable0RTT, + false, + nil, + utils.DefaultLogger, + protocol.Version1, + ) + require.NoError(t, err) + require.Nil(t, conn.testHooks) + conn.packer = packer + for _, opt := range opts { + opt(conn.Conn) + } + return &testConnection{ + conn: conn.Conn, + connRunner: connRunner, + sendConn: sendConn, + packer: packer, + destConnID: destConnID, + srcConnID: srcConnID, + } +} + +func TestConnectionHandleStreamRelatedFrames(t *testing.T) { + const id protocol.StreamID = 5 + connID := protocol.ConnectionID{} + + tests := []struct { + name string + frame wire.Frame + }{ + {name: "RESET_STREAM", frame: &wire.ResetStreamFrame{StreamID: id, ErrorCode: 42, FinalSize: 1337}}, + {name: "STOP_SENDING", frame: &wire.StopSendingFrame{StreamID: id, ErrorCode: 42}}, + {name: "MAX_STREAM_DATA", frame: &wire.MaxStreamDataFrame{StreamID: id, MaximumStreamData: 1337}}, + {name: "STREAM_DATA_BLOCKED", frame: &wire.StreamDataBlockedFrame{StreamID: id, MaximumStreamData: 42}}, + {name: "STREAM_FRAME", frame: &wire.StreamFrame{StreamID: id, Data: []byte{1, 2, 3, 4, 5, 6, 7, 8}, Offset: 1337}}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + tc := newServerTestConnection(t, gomock.NewController(t), nil, false) + data, err := test.frame.Append(nil, protocol.Version1) + require.NoError(t, err) + _, _, _, err = tc.conn.handleFrames(data, connID, protocol.Encryption1RTT, nil, monotime.Now()) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + }) + } +} + +func TestConnectionHandleConnectionFlowControlFrames(t *testing.T) { + mockCtrl := gomock.NewController(t) + connFC := newConnectionFlowController(0, 0, nil, utils.NewRTTStats(), utils.DefaultLogger) + require.Zero(t, connFC.SendWindowSize()) + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptConnFlowController(connFC)) + now := monotime.Now() + connID := protocol.ConnectionID{} + // MAX_DATA frame + _, err := tc.conn.handleFrame(&wire.MaxDataFrame{MaximumData: 1337}, protocol.Encryption1RTT, connID, now) + require.NoError(t, err) + require.Equal(t, protocol.ByteCount(1337), connFC.SendWindowSize()) + // DATA_BLOCKED frame + _, err = tc.conn.handleFrame(&wire.DataBlockedFrame{MaximumData: 1337}, protocol.Encryption1RTT, connID, now) + require.NoError(t, err) +} + +func TestConnectionServerInvalidFrames(t *testing.T) { + mockCtrl := gomock.NewController(t) + tc := newServerTestConnection(t, mockCtrl, nil, false) + + for _, test := range []struct { + Name string + Frame wire.Frame + }{ + {Name: "NEW_TOKEN", Frame: &wire.NewTokenFrame{Token: []byte("foobar")}}, + {Name: "HANDSHAKE_DONE", Frame: &wire.HandshakeDoneFrame{}}, + {Name: "PATH_RESPONSE", Frame: &wire.PathResponseFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}}, + } { + t.Run(test.Name, func(t *testing.T) { + _, err := tc.conn.handleFrame(test.Frame, protocol.Encryption1RTT, protocol.ConnectionID{}, monotime.Now()) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + }) + } +} + +func TestConnectionClose(t *testing.T) { + t.Run("transport error", func(t *testing.T) { + expectedErr := &qerr.TransportError{ + ErrorCode: 1337, + FrameType: 42, + ErrorMessage: "foobar", + } + testConnectionClose(t, false, expectedErr) + }) + t.Run("application error", func(t *testing.T) { + expectedErr := &qerr.ApplicationError{ + ErrorCode: 1337, + ErrorMessage: "foobar", + } + testConnectionClose(t, true, expectedErr) + }) +} + +func testConnectionClose(t *testing.T, useApplicationClose bool, expectedErr error) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + errChan := make(chan error, 1) + + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + b := getPacketBuffer() + b.Data = append(b.Data, []byte("connection close")...) + if useApplicationClose { + tc.packer.EXPECT().PackApplicationClose(expectedErr, gomock.Any(), protocol.Version1).Return(&coalescedPacket{buffer: b}, nil) + } else { + tc.packer.EXPECT().PackConnectionClose(expectedErr, gomock.Any(), protocol.Version1).Return(&coalescedPacket{buffer: b}, nil) + } + tc.sendConn.EXPECT().Write([]byte("connection close"), gomock.Any(), gomock.Any()) + tc.connRunner.EXPECT().ReplaceWithClosed(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + go func() { errChan <- tc.conn.run() }() + tc.conn.closeLocal(expectedErr) + + synctest.Wait() + + var want qlog.ConnectionClosed + if useApplicationClose { + code := expectedErr.(*qerr.ApplicationError).ErrorCode + want = qlog.ConnectionClosed{ + Initiator: qlog.InitiatorLocal, + ApplicationError: &code, + Reason: expectedErr.(*qerr.ApplicationError).ErrorMessage, + } + } else { + code := expectedErr.(*qerr.TransportError).ErrorCode + want = qlog.ConnectionClosed{ + Initiator: qlog.InitiatorLocal, + ConnectionError: &code, + Reason: expectedErr.(*qerr.TransportError).ErrorMessage, + } + } + require.Equal(t, + []qlogwriter.Event{want}, + eventRecorder.Events(qlog.ConnectionClosed{}), + ) + eventRecorder.Clear() + + select { + case err := <-errChan: + require.ErrorIs(t, err, expectedErr) + default: + t.Fatal("connection was not closed") + } + + // further calls to CloseWithError don't do anything + tc.conn.CloseWithError(42, "another error") + require.Empty(t, eventRecorder.Events(qlog.ConnectionClosed{})) + }) +} + +func TestConnectionStatelessReset(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + errChan := make(chan error, 1) + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + + go func() { errChan <- tc.conn.run() }() + tc.conn.destroy(&StatelessResetError{}) + + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{qlog.ConnectionClosed{Initiator: qlog.InitiatorLocal, Trigger: qlog.ConnectionCloseTriggerStatelessReset}}, + eventRecorder.Events(qlog.ConnectionClosed{}), + ) + }) +} + +func getLongHeaderPacket(t *testing.T, remoteAddr net.Addr, extHdr *wire.ExtendedHeader, data []byte) receivedPacket { + t.Helper() + b, err := extHdr.Append(nil, protocol.Version1) + require.NoError(t, err) + return receivedPacket{ + remoteAddr: remoteAddr, + data: append(b, data...), + buffer: getPacketBuffer(), + rcvTime: monotime.Now(), + } +} + +func getShortHeaderPacket(t *testing.T, remoteAddr net.Addr, connID protocol.ConnectionID, pn protocol.PacketNumber, data []byte) receivedPacket { + t.Helper() + b, err := wire.AppendShortHeader(nil, connID, pn, protocol.PacketNumberLen2, protocol.KeyPhaseOne) + require.NoError(t, err) + return receivedPacket{ + remoteAddr: remoteAddr, + data: append(b, data...), + buffer: getPacketBuffer(), + rcvTime: monotime.Now(), + } +} + +func TestConnectionServerInvalidPackets(t *testing.T) { + t.Run("Retry", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + p := getLongHeaderPacket(t, + tc.remoteAddr, + &wire.ExtendedHeader{Header: wire.Header{ + Type: protocol.PacketTypeRetry, + DestConnectionID: tc.conn.origDestConnID, + SrcConnectionID: tc.srcConnID, + Version: tc.conn.version, + Token: []byte("foobar"), + }}, + make([]byte, 16), /* Retry integrity tag */ + ) + wasProcessed, err := tc.conn.handleOnePacket(p, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: tc.srcConnID, + DestConnectionID: tc.conn.origDestConnID, + Version: tc.conn.version, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + }) + + t.Run("version negotiation", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + b := wire.ComposeVersionNegotiation( + protocol.ArbitraryLenConnectionID(tc.srcConnID.Bytes()), + protocol.ArbitraryLenConnectionID(tc.conn.origDestConnID.Bytes()), + []Version{Version1}, + ) + wasProcessed, err := tc.conn.handleOnePacket(receivedPacket{data: b, buffer: getPacketBuffer()}, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: len(b)}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + }) + + t.Run("unsupported version", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + p := getLongHeaderPacket(t, + tc.remoteAddr, + &wire.ExtendedHeader{ + Header: wire.Header{Type: protocol.PacketTypeHandshake, Version: 1234}, + PacketNumberLen: protocol.PacketNumberLen2, + }, + nil, + ) + wasProcessed, err := tc.conn.handleOnePacket(p, 42) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{Version: 1234}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: 42, + Trigger: qlog.PacketDropUnsupportedVersion, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + }) + + t.Run("invalid header", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + p := getLongHeaderPacket(t, + tc.remoteAddr, + &wire.ExtendedHeader{ + Header: wire.Header{Type: protocol.PacketTypeHandshake, Version: Version1}, + PacketNumberLen: protocol.PacketNumberLen2, + }, + nil, + ) + p.data[0] ^= 0x40 // unset the QUIC bit + wasProcessed, err := tc.conn.handleOnePacket(p, 42) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: 42, + Trigger: qlog.PacketDropHeaderParseError, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + }) +} + +func TestConnectionClientDrop0RTT(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newClientTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + p := getLongHeaderPacket(t, + tc.remoteAddr, + &wire.ExtendedHeader{ + Header: wire.Header{Type: protocol.PacketType0RTT, Length: 2, Version: protocol.Version1}, + PacketNumberLen: protocol.PacketNumberLen2, + }, + nil, + ) + wasProcessed, err := tc.conn.handleOnePacket(p, 1234) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + DatagramPayloadChecksum: 1234, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) +} + +func TestConnectionUnpacking(t *testing.T) { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptUnpacker(unpacker), + connectionOptTracer(&eventRecorder), + ) + + // receive a long header packet + hdr := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: tc.srcConnID, + Version: protocol.Version1, + Length: 1, + }, + PacketNumber: 0x37, + PacketNumberLen: protocol.PacketNumberLen1, + } + unpackedHdr := *hdr + unpackedHdr.PacketNumber = 0x1337 + packet := getLongHeaderPacket(t, tc.remoteAddr, hdr, nil) + packet.ecn = protocol.ECNCE + rcvTime := monotime.Now().Add(-10 * time.Second) + packet.rcvTime = rcvTime + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return(&unpackedPacket{ + encryptionLevel: protocol.EncryptionInitial, + hdr: &unpackedHdr, + data: []byte{0}, // one PADDING frame + }, nil) + + wasProcessed, err := tc.conn.handleOnePacket(packet, 42) + require.NoError(t, err) + require.True(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + DestConnectionID: tc.srcConnID, + PacketNumber: protocol.PacketNumber(0x1337), + Version: protocol.Version1, + }, + Frames: []qlog.Frame{}, + ECN: qlog.ECNCE, + Raw: qlog.RawInfo{Length: int(packet.Size()), PayloadLength: 1}, + DatagramPayloadChecksum: 42, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + // receive a duplicate of this packet + packet = getLongHeaderPacket(t, tc.remoteAddr, hdr, nil) + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return(&unpackedPacket{ + encryptionLevel: protocol.EncryptionInitial, + hdr: &unpackedHdr, + data: []byte{0}, // one PADDING frame + }, nil) + wasProcessed, err = tc.conn.handleOnePacket(packet, 43) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + DestConnectionID: tc.srcConnID, + PacketNumber: protocol.PacketNumber(0x1337), + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(packet.Size()), PayloadLength: 1}, + DatagramPayloadChecksum: 43, + Trigger: qlog.PacketDropDuplicate, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + // receive a short header packet + packet = getShortHeaderPacket(t, tc.remoteAddr, tc.srcConnID, 0x37, nil) + packet.ecn = protocol.ECT1 + packet.rcvTime = rcvTime + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return( + protocol.PacketNumber(0x1337), protocol.PacketNumberLen2, protocol.KeyPhaseZero, []byte{0} /* PADDING */, nil, + ) + wasProcessed, err = tc.conn.handleOnePacket(packet, 0) + require.NoError(t, err) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType1RTT, + DestConnectionID: tc.srcConnID, + PacketNumber: protocol.PacketNumber(0x1337), + KeyPhaseBit: protocol.KeyPhaseZero, + }, + Raw: qlog.RawInfo{Length: int(packet.Size())}, + Frames: []qlog.Frame{}, + ECN: qlog.ECT1, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketDropped{}), + ) + require.True(t, wasProcessed) +} + +func TestConnectionUnpackCoalescedPacket(t *testing.T) { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptUnpacker(unpacker), + connectionOptTracer(&eventRecorder), + ) + hdr1 := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: tc.srcConnID, + Version: protocol.Version1, + Length: 1, + }, + PacketNumber: 37, + PacketNumberLen: protocol.PacketNumberLen1, + } + hdr2 := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: tc.srcConnID, + Version: protocol.Version1, + Length: 1, + }, + PacketNumber: 38, + PacketNumberLen: protocol.PacketNumberLen1, + } + // add a packet with a different source connection ID + incorrectSrcConnID := protocol.ParseConnectionID([]byte{0xa, 0xb, 0xc}) + hdr3 := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: incorrectSrcConnID, + Version: protocol.Version1, + Length: 1, + }, + PacketNumber: 0x42, + PacketNumberLen: protocol.PacketNumberLen1, + } + unpackedHdr1 := *hdr1 + unpackedHdr1.PacketNumber = 1337 + unpackedHdr2 := *hdr2 + unpackedHdr2.PacketNumber = 1338 + + packet := getLongHeaderPacket(t, tc.remoteAddr, hdr1, nil) + firstPacketLen := packet.Size() + packet2 := getLongHeaderPacket(t, tc.remoteAddr, hdr2, nil) + packet3 := getLongHeaderPacket(t, tc.remoteAddr, hdr3, nil) + packet.data = append(packet.data, packet2.data...) + packet.data = append(packet.data, packet3.data...) + packet.ecn = protocol.ECT1 + rcvTime := monotime.Now() + packet.rcvTime = rcvTime + + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return(&unpackedPacket{ + encryptionLevel: protocol.EncryptionInitial, + hdr: &unpackedHdr1, + data: []byte{0}, // one PADDING frame + }, nil) + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return(&unpackedPacket{ + encryptionLevel: protocol.EncryptionHandshake, + hdr: &unpackedHdr2, + data: []byte{1}, // one PING frame + }, nil) + wasProcessed, err := tc.conn.handleOnePacket(packet, 42) + require.NoError(t, err) + require.True(t, wasProcessed) + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + DestConnectionID: tc.srcConnID, + PacketNumber: protocol.PacketNumber(1337), + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(firstPacketLen), PayloadLength: 1}, + DatagramPayloadChecksum: 42, + Frames: []qlog.Frame{}, + ECN: qlog.ECT1, + }, + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + DestConnectionID: tc.srcConnID, + PacketNumber: protocol.PacketNumber(1338), + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(packet2.Size()), PayloadLength: 1}, + DatagramPayloadChecksum: 42, + Frames: []qlog.Frame{{Frame: &wire.PingFrame{}}}, + ECN: qlog.ECT1, + }, + qlog.PacketDropped{ + Header: qlog.PacketHeader{DestConnectionID: incorrectSrcConnID}, + Raw: qlog.RawInfo{Length: int(packet3.Size())}, + DatagramPayloadChecksum: 42, + Trigger: qlog.PacketDropUnknownConnectionID, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketDropped{}), + ) +} + +func TestConnectionUnpackFailuresFatal(t *testing.T) { + t.Run("other errors", func(t *testing.T) { + require.ErrorIs(t, + testConnectionUnpackFailureFatal(t, &qerr.TransportError{ErrorCode: qerr.ConnectionIDLimitError}), + &qerr.TransportError{ErrorCode: qerr.ConnectionIDLimitError}, + ) + }) + + t.Run("invalid reserved bits", func(t *testing.T) { + require.ErrorIs(t, + testConnectionUnpackFailureFatal(t, wire.ErrInvalidReservedBits), + &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}, + ) + }) +} + +func testConnectionUnpackFailureFatal(t *testing.T, unpackErr error) error { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptUnpacker(unpacker), + ) + + tc.connRunner.EXPECT().ReplaceWithClosed(gomock.Any(), gomock.Any(), gomock.Any()) + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(0), protocol.PacketNumberLen(0), protocol.KeyPhaseBit(0), nil, unpackErr) + tc.packer.EXPECT().PackConnectionClose(gomock.Any(), gomock.Any(), protocol.Version1).Return(&coalescedPacket{buffer: getPacketBuffer()}, nil) + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()) + tc.conn.handlePacket(getShortHeaderPacket(t, tc.remoteAddr, tc.srcConnID, 0x42, nil)) + + select { + case err := <-errChan: + require.Error(t, err) + return err + case <-time.After(time.Second): + t.Fatal("timeout") + } + return nil +} + +func TestConnectionUnpackFailureDropped(t *testing.T) { + t.Run("keys dropped", func(t *testing.T) { + testConnectionUnpackFailureDropped(t, handshake.ErrKeysDropped, qlog.PacketDropKeyUnavailable) + }) + + t.Run("decryption failed", func(t *testing.T) { + testConnectionUnpackFailureDropped(t, handshake.ErrDecryptionFailed, qlog.PacketDropPayloadDecryptError) + }) + + t.Run("header parse error", func(t *testing.T) { + testConnectionUnpackFailureDropped(t, &headerParseError{err: assert.AnError}, qlog.PacketDropHeaderParseError) + }) +} + +func testConnectionUnpackFailureDropped(t *testing.T, unpackErr error, packetDropReason qlog.PacketDropReason) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptUnpacker(unpacker), + connectionOptTracer(&eventRecorder), + ) + + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(0), protocol.PacketNumberLen(0), protocol.KeyPhaseBit(0), nil, unpackErr) + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + packet := getShortHeaderPacket(t, tc.remoteAddr, tc.srcConnID, 0x42, nil) + tc.conn.handlePacket(packet) + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType1RTT, + DestConnectionID: tc.srcConnID, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(packet.Size())}, + Trigger: packetDropReason, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case <-errChan: + default: + t.Fatal("timeout") + } + }) +} + +func TestConnectionMaxUnprocessedPackets(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + for range protocol.MaxConnUnprocessedPackets { + // nothing here should block + tc.conn.handlePacket(receivedPacket{data: []byte("foobar")}) + } + tc.conn.handlePacket(receivedPacket{data: []byte("foobar")}) + + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: 6}, + Trigger: qlog.PacketDropDOSPrevention, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + }) +} + +func TestConnectionRemoteClose(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + unpacker := NewMockUnpacker(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + connectionOptUnpacker(unpacker), + ) + ccf, err := (&wire.ConnectionCloseFrame{ + ErrorCode: uint64(qerr.StreamLimitError), + ReasonPhrase: "foobar", + }).Append(nil, protocol.Version1) + require.NoError(t, err) + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(1), protocol.PacketNumberLen2, protocol.KeyPhaseBit(0), ccf, nil) + + tc.connRunner.EXPECT().ReplaceWithClosed(gomock.Any(), gomock.Any(), gomock.Any()) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + p := getShortHeaderPacket(t, tc.remoteAddr, tc.srcConnID, 1, []byte("encrypted")) + tc.conn.handlePacket(receivedPacket{data: p.data, buffer: p.buffer, rcvTime: monotime.Now()}) + + synctest.Wait() + + expectedErr := &qerr.TransportError{ErrorCode: qerr.StreamLimitError, ErrorMessage: "foobar", Remote: true} + select { + case err := <-errChan: + require.ErrorIs(t, err, expectedErr) + default: + t.Fatal("timeout") + } + + code := expectedErr.ErrorCode + require.Equal(t, + []qlogwriter.Event{ + qlog.ConnectionClosed{ + Initiator: qlog.InitiatorRemote, + ConnectionError: &code, + Reason: expectedErr.ErrorMessage, + }, + }, + eventRecorder.Events(qlog.ConnectionClosed{}), + ) + }) +} + +func TestConnectionIdleTimeoutDuringHandshake(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const timeout = 7 * time.Second + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, + mockCtrl, + &Config{HandshakeIdleTimeout: timeout}, + false, + connectionOptTracer(&eventRecorder), + ) + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).AnyTimes() + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + start := monotime.Now() + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, &IdleTimeoutError{}) + require.Equal(t, timeout, monotime.Since(start)) + case <-time.After(timeout + time.Nanosecond): + t.Fatal("timeout") + } + + require.Equal(t, + []qlogwriter.Event{ + qlog.ConnectionClosed{ + Initiator: qlog.InitiatorLocal, + Trigger: qlog.ConnectionCloseTriggerIdleTimeout, + }, + }, + eventRecorder.Events(qlog.ConnectionClosed{}), + ) + }) +} + +func TestConnectionHandshakeIdleTimeout(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, + mockCtrl, + &Config{HandshakeIdleTimeout: 7 * time.Second}, + false, + connectionOptTracer(&eventRecorder), + func(c *Conn) { c.creationTime = monotime.Now().Add(-20 * time.Second) }, + ) + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).AnyTimes() + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, &HandshakeTimeoutError{}) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Equal(t, + []qlogwriter.Event{ + qlog.ConnectionClosed{ + Initiator: qlog.InitiatorLocal, + Trigger: qlog.ConnectionCloseTriggerIdleTimeout, + }, + }, + eventRecorder.Events(qlog.ConnectionClosed{}), + ) + }) +} + +func TestConnectionTransportParameters(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + connFC := newConnectionFlowController(0, 0, nil, utils.NewRTTStats(), utils.DefaultLogger) + require.Zero(t, connFC.SendWindowSize()) + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + connectionOptConnFlowController(connFC), + ) + _, err := tc.conn.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + _, err = tc.conn.OpenUniStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + params := &wire.TransportParameters{ + MaxIdleTimeout: 90 * time.Second, + InitialMaxStreamDataBidiLocal: 0x5000, + InitialMaxData: 1337, + ActiveConnectionIDLimit: 3, + // marshaling always sets it to this value + MaxUDPPayloadSize: protocol.MaxPacketBufferSize, + OriginalDestinationConnectionID: tc.destConnID, + MaxBidiStreamNum: 1, + MaxUniStreamNum: 1, + } + require.NoError(t, tc.conn.handleTransportParameters(params)) + require.Equal(t, protocol.ByteCount(1337), connFC.SendWindowSize()) + _, err = tc.conn.OpenStream() + require.NoError(t, err) + _, err = tc.conn.OpenUniStream() + require.NoError(t, err) + + require.Equal(t, + []qlogwriter.Event{ + qlog.ParametersSet{ + Initiator: qlog.InitiatorRemote, + MaxIdleTimeout: 90 * time.Second, + InitialMaxStreamDataBidiLocal: 0x5000, + InitialMaxData: 1337, + ActiveConnectionIDLimit: 3, + // marshaling always sets it to this value + MaxUDPPayloadSize: protocol.MaxPacketBufferSize, + OriginalDestinationConnectionID: tc.destConnID, + InitialMaxStreamsBidi: 1, + InitialMaxStreamsUni: 1, + }, + }, + eventRecorder.Events(qlog.ParametersSet{}), + ) +} + +func TestConnectionHandleMaxStreamsFrame(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + connFC := newConnectionFlowController(0, 0, nil, utils.NewRTTStats(), utils.DefaultLogger) + tc := newServerTestConnection(t, mockCtrl, nil, false, connectionOptConnFlowController(connFC)) + tc.conn.handleTransportParameters(&wire.TransportParameters{}) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + uniStreamChan := make(chan error) + go func() { + _, err := tc.conn.OpenUniStreamSync(ctx) + uniStreamChan <- err + }() + bidiStreamChan := make(chan error) + go func() { + _, err := tc.conn.OpenStreamSync(ctx) + bidiStreamChan <- err + }() + + synctest.Wait() + select { + case <-uniStreamChan: + t.Fatal("uni stream should be blocked") + case <-bidiStreamChan: + t.Fatal("bidi stream should be blocked") + default: + } + + // MAX_STREAMS frame for bidirectional stream + _, err := tc.conn.handleFrame( + &wire.MaxStreamsFrame{Type: protocol.StreamTypeBidi, MaxStreamNum: 10}, + protocol.Encryption1RTT, + protocol.ConnectionID{}, + monotime.Now(), + ) + require.NoError(t, err) + + synctest.Wait() + + select { + case <-uniStreamChan: + t.Fatal("uni stream should be blocked") + default: + } + select { + case err := <-bidiStreamChan: + require.NoError(t, err) + default: + t.Fatal("bidi stream should be unblocked") + } + + // MAX_STREAMS frame for bidirectional stream + _, err = tc.conn.handleFrame( + &wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: 10}, + protocol.Encryption1RTT, + protocol.ConnectionID{}, + monotime.Now(), + ) + require.NoError(t, err) + + synctest.Wait() + select { + case err := <-uniStreamChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + }) +} + +func TestConnectionTransportParameterValidationFailureServer(t *testing.T) { + tc := newServerTestConnection(t, nil, nil, false) + err := tc.conn.handleTransportParameters(&wire.TransportParameters{ + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + }) + assert.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.TransportParameterError}) + assert.ErrorContains(t, err, "expected initial_source_connection_id to equal") +} + +func TestConnectionTransportParameterValidationFailureClient(t *testing.T) { + t.Run("initial_source_connection_id", func(t *testing.T) { + tc := newClientTestConnection(t, nil, nil, false) + err := tc.conn.handleTransportParameters(&wire.TransportParameters{ + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + }) + assert.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.TransportParameterError}) + assert.ErrorContains(t, err, "expected initial_source_connection_id to equal") + }) + + t.Run("original_destination_connection_id", func(t *testing.T) { + tc := newClientTestConnection(t, nil, nil, false) + err := tc.conn.handleTransportParameters(&wire.TransportParameters{ + InitialSourceConnectionID: tc.destConnID, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + }) + assert.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.TransportParameterError}) + assert.ErrorContains(t, err, "expected original_destination_connection_id to equal") + }) + + t.Run("retry_source_connection_id if no retry", func(t *testing.T) { + tc := newClientTestConnection(t, nil, nil, false) + rcid := protocol.ParseConnectionID([]byte{1, 2, 3, 4}) + params := &wire.TransportParameters{ + InitialSourceConnectionID: tc.destConnID, + OriginalDestinationConnectionID: tc.destConnID, + RetrySourceConnectionID: &rcid, + } + err := tc.conn.handleTransportParameters(params) + assert.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.TransportParameterError}) + assert.ErrorContains(t, err, "received retry_source_connection_id, although no Retry was performed") + }) + + t.Run("retry_source_connection_id missing", func(t *testing.T) { + tc := newClientTestConnection(t, + nil, + nil, + false, + connectionOptRetrySrcConnID(protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef})), + ) + params := &wire.TransportParameters{ + InitialSourceConnectionID: tc.destConnID, + OriginalDestinationConnectionID: tc.destConnID, + } + err := tc.conn.handleTransportParameters(params) + assert.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.TransportParameterError}) + assert.ErrorContains(t, err, "missing retry_source_connection_id") + }) + + t.Run("retry_source_connection_id incorrect", func(t *testing.T) { + tc := newClientTestConnection(t, + nil, + nil, + false, + connectionOptRetrySrcConnID(protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef})), + ) + wrongCID := protocol.ParseConnectionID([]byte{1, 2, 3, 4}) + params := &wire.TransportParameters{ + InitialSourceConnectionID: tc.destConnID, + OriginalDestinationConnectionID: tc.destConnID, + RetrySourceConnectionID: &wrongCID, + } + err := tc.conn.handleTransportParameters(params) + assert.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.TransportParameterError}) + assert.ErrorContains(t, err, "expected retry_source_connection_id to equal") + }) +} + +func TestConnectionHandshakeServer(t *testing.T) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := NewMockUnpacker(mockCtrl) + tc := newServerTestConnection( + t, + mockCtrl, + nil, + false, + connectionOptCryptoSetup(cs), + connectionOptUnpacker(unpacker), + ) + + // the state transition is driven by processing of a CRYPTO frame + hdr := &wire.ExtendedHeader{ + Header: wire.Header{Type: protocol.PacketTypeHandshake, Version: protocol.Version1}, + PacketNumberLen: protocol.PacketNumberLen2, + } + data, err := (&wire.CryptoFrame{Data: []byte("foobar")}).Append(nil, protocol.Version1) + require.NoError(t, err) + + cs.EXPECT().DiscardInitialKeys().Times(2) + gomock.InOrder( + cs.EXPECT().StartHandshake(gomock.Any()), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}), + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return( + &unpackedPacket{hdr: hdr, encryptionLevel: protocol.EncryptionHandshake, data: data}, nil, + ), + cs.EXPECT().HandleMessage([]byte("foobar"), protocol.EncryptionHandshake), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventHandshakeComplete}), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}), + cs.EXPECT().SetHandshakeConfirmed(), + cs.EXPECT().GetSessionTicket().Return([]byte("session ticket"), nil), + ) + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack).AnyTimes() + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + p := getLongHeaderPacket(t, tc.remoteAddr, hdr, nil) + tc.conn.handlePacket(receivedPacket{data: p.data, buffer: p.buffer, rcvTime: monotime.Now()}) + + select { + case <-tc.conn.HandshakeComplete(): + case <-tc.conn.Context().Done(): + t.Fatal("connection context done") + case <-time.After(time.Second): + t.Fatal("timeout") + } + + var foundSessionTicket, foundHandshakeDone, foundNewToken bool + frames, _, _ := tc.conn.framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + for _, frame := range frames { + switch f := frame.Frame.(type) { + case *wire.CryptoFrame: + assert.Equal(t, []byte("session ticket"), f.Data) + foundSessionTicket = true + case *wire.HandshakeDoneFrame: + foundHandshakeDone = true + case *wire.NewTokenFrame: + assert.NotEmpty(t, f.Token) + foundNewToken = true + } + } + assert.True(t, foundSessionTicket) + assert.True(t, foundHandshakeDone) + assert.True(t, foundNewToken) + + // test teardown + cs.EXPECT().Close() + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnectionHandshakeClient(t *testing.T) { + t.Run("without preferred address", func(t *testing.T) { + testConnectionHandshakeClient(t, false) + }) + t.Run("with preferred address", func(t *testing.T) { + testConnectionHandshakeClient(t, true) + }) +} + +func testConnectionHandshakeClient(t *testing.T, usePreferredAddress bool) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := NewMockUnpacker(mockCtrl) + tc := newClientTestConnection(t, mockCtrl, nil, false, connectionOptCryptoSetup(cs), connectionOptUnpacker(unpacker)) + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + // the state transition is driven by processing of a CRYPTO frame + hdr := &wire.ExtendedHeader{ + Header: wire.Header{Type: protocol.PacketTypeHandshake, Version: protocol.Version1}, + PacketNumberLen: protocol.PacketNumberLen2, + } + data, err := (&wire.CryptoFrame{Data: []byte("foobar")}).Append(nil, protocol.Version1) + require.NoError(t, err) + + tp := &wire.TransportParameters{ + OriginalDestinationConnectionID: tc.destConnID, + MaxIdleTimeout: time.Hour, + } + preferredAddressConnID := protocol.ParseConnectionID([]byte{10, 8, 6, 4}) + preferredAddressResetToken := protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1} + if usePreferredAddress { + tp.PreferredAddress = &wire.PreferredAddress{ + IPv4: netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 42), + IPv6: netip.AddrPortFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}), 13), + ConnectionID: preferredAddressConnID, + StatelessResetToken: preferredAddressResetToken, + } + } + + packedFirstPacket := make(chan struct{}) + gomock.InOrder( + cs.EXPECT().StartHandshake(gomock.Any()), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}), + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).DoAndReturn( + func(b bool, bc protocol.ByteCount, t monotime.Time, v protocol.Version) (*coalescedPacket, error) { + close(packedFirstPacket) + return &coalescedPacket{buffer: getPacketBuffer(), longHdrPackets: []*longHeaderPacket{{header: hdr}}}, nil + }, + ), + // initial keys are dropped when the first handshake packet is sent + cs.EXPECT().DiscardInitialKeys(), + // no more data to send + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return( + &unpackedPacket{hdr: hdr, encryptionLevel: protocol.EncryptionHandshake, data: data}, nil, + ), + cs.EXPECT().HandleMessage([]byte("foobar"), protocol.EncryptionHandshake), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventReceivedTransportParameters, TransportParameters: tp}), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventHandshakeComplete}), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}), + ) + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).Return(nil, nil).AnyTimes() + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + select { + case <-packedFirstPacket: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + p := getLongHeaderPacket(t, tc.remoteAddr, hdr, nil) + tc.conn.handlePacket(receivedPacket{data: p.data, buffer: p.buffer, rcvTime: monotime.Now()}) + + select { + case <-tc.conn.HandshakeComplete(): + case <-tc.conn.Context().Done(): + t.Fatal("connection context done") + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.True(t, mockCtrl.Satisfied()) + // the handshake isn't confirmed until we receive a HANDSHAKE_DONE frame from the server + + data, err = (&wire.HandshakeDoneFrame{}).Append(nil, protocol.Version1) + require.NoError(t, err) + done := make(chan struct{}) + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).Return(nil, nil).AnyTimes() + gomock.InOrder( + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return( + &unpackedPacket{hdr: hdr, encryptionLevel: protocol.Encryption1RTT, data: data}, nil, + ), + cs.EXPECT().DiscardInitialKeys(), + cs.EXPECT().SetHandshakeConfirmed(), + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buf *packetBuffer, _ protocol.ByteCount, _ monotime.Time, _ protocol.Version) (shortHeaderPacket, error) { + close(done) + return shortHeaderPacket{}, errNothingToPack + }, + ), + ) + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack).AnyTimes() + p = getLongHeaderPacket(t, tc.remoteAddr, hdr, nil) + tc.conn.handlePacket(receivedPacket{data: p.data, buffer: p.buffer, rcvTime: monotime.Now()}) + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + if usePreferredAddress { + tc.connRunner.EXPECT().AddResetToken(preferredAddressResetToken, gomock.Any()) + } + nextConnID := tc.conn.connIDManager.Get() + if usePreferredAddress { + require.Equal(t, preferredAddressConnID, nextConnID) + } + + // test teardown + cs.EXPECT().Close() + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + if usePreferredAddress { + tc.connRunner.EXPECT().RemoveResetToken(preferredAddressResetToken) + } + tc.conn.destroy(nil) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnection0RTTTransportParameters(t *testing.T) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := NewMockUnpacker(mockCtrl) + tc := newClientTestConnection(t, mockCtrl, nil, false, connectionOptCryptoSetup(cs), connectionOptUnpacker(unpacker)) + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + // the state transition is driven by processing of a CRYPTO frame + hdr := &wire.ExtendedHeader{ + Header: wire.Header{Type: protocol.PacketTypeHandshake, Version: protocol.Version1}, + PacketNumberLen: protocol.PacketNumberLen2, + } + data, err := (&wire.CryptoFrame{Data: []byte("foobar")}).Append(nil, protocol.Version1) + require.NoError(t, err) + + restored := &wire.TransportParameters{ + ActiveConnectionIDLimit: 3, + InitialMaxData: 0x5000, + InitialMaxStreamDataBidiLocal: 0x5000, + InitialMaxStreamDataBidiRemote: 1000, + InitialMaxStreamDataUni: 1000, + MaxBidiStreamNum: 500, + MaxUniStreamNum: 500, + } + new := *restored + new.MaxBidiStreamNum-- // the server is not allowed to reduce the limit + new.OriginalDestinationConnectionID = tc.destConnID + + packedFirstPacket := make(chan struct{}) + gomock.InOrder( + cs.EXPECT().StartHandshake(gomock.Any()), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventRestoredTransportParameters, TransportParameters: restored}), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}), + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).DoAndReturn( + func(b bool, bc protocol.ByteCount, t monotime.Time, v protocol.Version) (*coalescedPacket, error) { + close(packedFirstPacket) + return &coalescedPacket{buffer: getPacketBuffer(), longHdrPackets: []*longHeaderPacket{{header: hdr}}}, nil + }, + ), + // initial keys are dropped when the first handshake packet is sent + cs.EXPECT().DiscardInitialKeys(), + // no more data to send + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return( + &unpackedPacket{hdr: hdr, encryptionLevel: protocol.EncryptionHandshake, data: data}, nil, + ), + cs.EXPECT().HandleMessage([]byte("foobar"), protocol.EncryptionHandshake), + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventReceivedTransportParameters, TransportParameters: &new}), + cs.EXPECT().ConnectionState().Return(handshake.ConnectionState{Used0RTT: true}), + // cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}), + cs.EXPECT().Close(), + ) + tc.packer.EXPECT().PackCoalescedPacket(false, gomock.Any(), gomock.Any(), protocol.Version1).Return(nil, nil).AnyTimes() + tc.packer.EXPECT().PackConnectionClose(gomock.Any(), gomock.Any(), protocol.Version1).Return(&coalescedPacket{buffer: getPacketBuffer()}, nil) + tc.connRunner.EXPECT().ReplaceWithClosed(gomock.Any(), gomock.Any(), gomock.Any()) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + select { + case <-packedFirstPacket: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + p := getLongHeaderPacket(t, tc.remoteAddr, hdr, nil) + tc.conn.handlePacket(receivedPacket{data: p.data, buffer: p.buffer, rcvTime: monotime.Now()}) + + select { + case err := <-errChan: + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + require.ErrorContains(t, err, "server sent reduced limits after accepting 0-RTT data") + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnectionReceivePrioritization(t *testing.T) { + for _, handshakeComplete := range []bool{true, false} { + t.Run(fmt.Sprintf("handshake complete: %t", handshakeComplete), func(t *testing.T) { + events := testConnectionReceivePrioritization(t, handshakeComplete, 5) + require.Equal(t, []string{"unpack", "unpack", "unpack", "unpack", "unpack", "pack"}, events) + }) + } +} + +func testConnectionReceivePrioritization(t *testing.T, handshakeComplete bool, numPackets int) []string { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + opts := []testConnectionOpt{connectionOptUnpacker(unpacker)} + if handshakeComplete { + opts = append(opts, connectionOptHandshakeConfirmed()) + } + tc := newServerTestConnection(t, mockCtrl, nil, false, opts...) + + var events []string + var counter int + var testDone bool + done := make(chan struct{}) + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).DoAndReturn( + func(rcvTime monotime.Time, data []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error) { + counter++ + if counter == numPackets { + testDone = true + } + events = append(events, "unpack") + return protocol.PacketNumber(counter), protocol.PacketNumberLen2, protocol.KeyPhaseZero, []byte{0, 1} /* PADDING, PING */, nil + }, + ).Times(numPackets) + switch handshakeComplete { + case false: + tc.packer.EXPECT().PackCoalescedPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(b bool, bc protocol.ByteCount, t monotime.Time, v protocol.Version) (*coalescedPacket, error) { + events = append(events, "pack") + if testDone { + close(done) + } + return nil, nil + }, + ).AnyTimes() + case true: + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(b *packetBuffer, bc protocol.ByteCount, t monotime.Time, v protocol.Version) (shortHeaderPacket, error) { + events = append(events, "pack") + if testDone { + close(done) + } + return shortHeaderPacket{}, errNothingToPack + }, + ).AnyTimes() + } + + for i := range numPackets { + tc.conn.handlePacket(getShortHeaderPacket(t, tc.remoteAddr, tc.srcConnID, protocol.PacketNumber(i), []byte("foobar"))) + } + + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } + return events +} + +func TestConnectionPacketBuffering(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + cs := mocks.NewMockCryptoSetup(mockCtrl) + var eventRecorder events.Recorder + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptUnpacker(unpacker), + connectionOptCryptoSetup(cs), + connectionOptTracer(&eventRecorder), + ) + + cs.EXPECT().DiscardInitialKeys() + + hdr1 := wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: tc.srcConnID, + SrcConnectionID: tc.destConnID, + Length: 8, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen1, + PacketNumber: 1, + } + hdr2 := hdr1 + hdr2.PacketNumber = 2 + cs.EXPECT().StartHandshake(gomock.Any()) + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}) + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return(nil, handshake.ErrKeysNotYetAvailable).Times(2) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + hdrs := make(map[string]*wire.ExtendedHeader) + + packet1 := getLongHeaderPacket(t, tc.remoteAddr, &hdr1, []byte("packet1")) + datagramPayloadChecksum1 := qlog.CalculateDatagramPayloadChecksum(packet1.data) + hdrs["packet1"] = &hdr1 + tc.conn.handlePacket(packet1) + packet2 := getLongHeaderPacket(t, tc.remoteAddr, &hdr2, []byte("packet2")) + datagramPayloadChecksum2 := qlog.CalculateDatagramPayloadChecksum(packet2.data) + hdrs["packet2"] = &hdr2 + tc.conn.handlePacket(packet2) + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketBuffered{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(packet1.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum1, + }, + qlog.PacketBuffered{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(packet2.Size())}, + DatagramPayloadChecksum: datagramPayloadChecksum2, + }, + }, + eventRecorder.Events(qlog.PacketBuffered{}), + ) + + eventRecorder.Clear() + + // Now send another packet. + // In reality, this packet would contain a CRYPTO frame that advances the TLS handshake + // such that new keys become available. + var packets []string + hdr3 := hdr1 + hdr3.PacketNumber = 3 + hdrs["packet3"] = &hdr3 + tc.packer.EXPECT().PackCoalescedPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes() + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventReceivedReadKeys}) + cs.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}) + + gomock.InOrder( + // packet 3 contains a CRYPTO frame and triggers the keys to become available + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).DoAndReturn( + func(hdr *wire.Header, data []byte) (*unpackedPacket, error) { + id := string(data[len(data)-7:]) + packets = append(packets, id) + cf := &wire.CryptoFrame{Data: []byte("foobar")} + b, _ := cf.Append(nil, protocol.Version1) + extHdr, ok := hdrs[id] + if !ok { + panic(fmt.Sprintf("unknown header: %v", id)) + } + return &unpackedPacket{hdr: extHdr, encryptionLevel: protocol.EncryptionHandshake, data: b}, nil + }, + ), + cs.EXPECT().HandleMessage(gomock.Any(), gomock.Any()), + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).DoAndReturn( + func(hdr *wire.Header, data []byte) (*unpackedPacket, error) { + id := string(data[len(data)-7:]) + extHdr, ok := hdrs[id] + if !ok { + panic(fmt.Sprintf("unknown header: %v", id)) + } + packets = append(packets, id) + return &unpackedPacket{hdr: extHdr, encryptionLevel: protocol.EncryptionHandshake, data: []byte{0} /* PADDING */}, nil + }, + ).Times(2), + ) + + packet3 := getLongHeaderPacket(t, tc.remoteAddr, &hdr3, []byte("packet3")) + datagramPayloadChecksum3 := qlog.CalculateDatagramPayloadChecksum(packet3.data) + tc.conn.handlePacket(packet3) + + synctest.Wait() + + // packet3 triggered the keys to become available + // packet1 and packet2 are processed from the buffer in order + require.Equal(t, []string{"packet3", "packet1", "packet2"}, packets) + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + DestConnectionID: tc.srcConnID, + SrcConnectionID: tc.destConnID, + PacketNumber: 3, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(packet3.Size()), PayloadLength: 8}, + DatagramPayloadChecksum: datagramPayloadChecksum3, + Frames: []qlog.Frame{{Frame: &qlog.CryptoFrame{Length: 6}}}, + }, + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + DestConnectionID: tc.srcConnID, + SrcConnectionID: tc.destConnID, + PacketNumber: 1, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(packet1.Size()), PayloadLength: 8}, + DatagramPayloadChecksum: datagramPayloadChecksum1, + Frames: []qlog.Frame{}, + }, + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + DestConnectionID: tc.srcConnID, + SrcConnectionID: tc.destConnID, + PacketNumber: 2, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(packet1.Size()), PayloadLength: 8}, + DatagramPayloadChecksum: datagramPayloadChecksum2, + Frames: []qlog.Frame{}, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketBuffered{}), + ) + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + cs.EXPECT().Close() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } + }) +} + +func TestConnectionPacketPacing(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + sender := NewMockSender(mockCtrl) + + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptSentPacketHandler(sph), + connectionOptSender(sender), + connectionOptHandshakeConfirmed(), + ) + sender.EXPECT().Run() + + const step = 50 * time.Millisecond + + sph.EXPECT().GetLossDetectionTimeout().Return(monotime.Now().Add(time.Hour)).AnyTimes() + gomock.InOrder( + // 1. allow 2 packets to be sent + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny), + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()), + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny), + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()), + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendPacingLimited), + // 2. become pacing limited for 25ms + sph.EXPECT().TimeUntilSend().DoAndReturn(func() monotime.Time { return monotime.Now().Add(step) }), + // 3. send another packet + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny), + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()), + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendPacingLimited), + // 4. become pacing limited for 25ms... + sph.EXPECT().TimeUntilSend().DoAndReturn(func() monotime.Time { return monotime.Now().Add(step) }), + // ... but this time we're still pacing limited when waking up. + // In this case, we can only send an ACK. + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendPacingLimited), + // 5. stop the test by becoming pacing limited forever + sph.EXPECT().TimeUntilSend().Return(monotime.Now().Add(time.Hour)), + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()), + ) + sph.EXPECT().ECNMode(gomock.Any()).AnyTimes() + for i := range 3 { + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), Version1).DoAndReturn( + func(buf *packetBuffer, _ protocol.ByteCount, _ monotime.Time, _ protocol.Version) (shortHeaderPacket, error) { + buf.Data = append(buf.Data, []byte("packet"+strconv.Itoa(i+1))...) + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(i + 1)}, nil + }, + ) + } + tc.packer.EXPECT().PackAckOnlyPacket(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(_ protocol.ByteCount, _ monotime.Time, _ protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + buf := getPacketBuffer() + buf.Data = []byte("ack") + return shortHeaderPacket{PacketNumber: 1}, buf, nil + }, + ) + sender.EXPECT().WouldBlock().AnyTimes() + + type sentPacket struct { + time monotime.Time + data []byte + } + sendChan := make(chan sentPacket, 10) + sender.EXPECT().Send(gomock.Any(), gomock.Any(), gomock.Any()).Do(func(b *packetBuffer, _ uint16, _ protocol.ECN) { + sendChan <- sentPacket{time: monotime.Now(), data: b.Data} + }).Times(4) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + var times []monotime.Time + for i := range 3 { + select { + case b := <-sendChan: + require.Equal(t, []byte("packet"+strconv.Itoa(i+1)), b.data) + times = append(times, b.time) + case <-time.After(time.Hour): + t.Fatal("should have sent a packet") + } + } + select { + case b := <-sendChan: + require.Equal(t, []byte("ack"), b.data) + times = append(times, b.time) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Equal(t, times[0], times[1]) + require.Equal(t, times[2], times[1].Add(step)) + require.Equal(t, times[3], times[2].Add(step)) + + synctest.Wait() // make sure that no more packets are sent + require.True(t, mockCtrl.Satisfied()) + + // test teardown + sender.EXPECT().Close() + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case <-sendChan: + t.Fatal("should not have sent any more packets") + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + }) +} + +// When the send queue blocks, we need to reset the pacing timer, otherwise the run loop might busy-loop. +// See https://github.com/apernet/quic-go/pull/4943 for more details. +func TestConnectionPacingAndSendQueue(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + sender := NewMockSender(mockCtrl) + + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptSentPacketHandler(sph), + connectionOptSender(sender), + connectionOptHandshakeConfirmed(), + ) + sender.EXPECT().Run() + + sendQueueAvailable := make(chan struct{}) + pacingDeadline := monotime.Now().Add(-time.Millisecond) + var counter int + // allow exactly one packet to be sent, then become blocked + sender.EXPECT().WouldBlock().Return(false) + sender.EXPECT().WouldBlock().DoAndReturn(func() bool { counter++; return true }).AnyTimes() + sender.EXPECT().Available().Return(sendQueueAvailable).AnyTimes() + sph.EXPECT().GetLossDetectionTimeout().Return(monotime.Now().Add(time.Hour)).AnyTimes() + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendPacingLimited).AnyTimes() + sph.EXPECT().TimeUntilSend().Return(pacingDeadline).AnyTimes() + sph.EXPECT().ECNMode(gomock.Any()).Return(protocol.ECNNon).AnyTimes() + tc.packer.EXPECT().PackAckOnlyPacket(gomock.Any(), gomock.Any(), gomock.Any()).Return( + shortHeaderPacket{}, nil, errNothingToPack, + ) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + sender.EXPECT().Close() + tc.conn.destroy(nil) + + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + + // make sure the run loop didn't do too many iterations + require.Less(t, counter, 3) + }) +} + +func TestConnectionIdleTimeout(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + &Config{MaxIdleTimeout: time.Minute}, + false, + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + connectionOptRTT(time.Millisecond), + ) + // the idle timeout is set when the transport parameters are received + const idleTimeout = 500 * time.Millisecond + require.NoError(t, tc.conn.handleTransportParameters(&wire.TransportParameters{ + MaxIdleTimeout: idleTimeout, + })) + + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).AnyTimes() + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()) + sph.EXPECT().ECNMode(gomock.Any()).AnyTimes() + var lastSendTime monotime.Time + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buf *packetBuffer, _ protocol.ByteCount, _ monotime.Time, _ protocol.Version) (shortHeaderPacket, error) { + buf.Data = append(buf.Data, []byte("foobar")...) + lastSendTime = monotime.Now() + return shortHeaderPacket{Frames: []ackhandler.Frame{{Frame: &wire.PingFrame{}}}, Length: 6}, nil + }, + ) + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack) + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()) + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, &IdleTimeoutError{}) + require.NotZero(t, lastSendTime) + require.Equal(t, idleTimeout, monotime.Since(lastSendTime)) + case <-time.After(time.Hour): + t.Fatal("should have timed out") + } + }) +} + +func TestConnectionKeepAlive(t *testing.T) { + t.Run("enabled", func(t *testing.T) { + testConnectionKeepAlive(t, true, true) + }) + + t.Run("disabled", func(t *testing.T) { + testConnectionKeepAlive(t, false, false) + }) +} + +func testConnectionKeepAlive(t *testing.T, enable, expectKeepAlive bool) { + synctest.Test(t, func(t *testing.T) { + var keepAlivePeriod time.Duration + if enable { + keepAlivePeriod = time.Second + } + + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + &Config{MaxIdleTimeout: time.Second, KeepAlivePeriod: keepAlivePeriod}, + false, + connectionOptUnpacker(unpacker), + connectionOptHandshakeConfirmed(), + connectionOptRTT(time.Millisecond), + ) + // the idle timeout is set when the transport parameters are received + const idleTimeout = 50 * time.Millisecond + require.NoError(t, tc.conn.handleTransportParameters(&wire.TransportParameters{ + MaxIdleTimeout: idleTimeout, + })) + + // Receive a packet. This starts the keep-alive timer. + buf := getPacketBuffer() + var err error + buf.Data, err = wire.AppendShortHeader(buf.Data, tc.srcConnID, 1, protocol.PacketNumberLen1, protocol.KeyPhaseZero) + require.NoError(t, err) + buf.Data = append(buf.Data, []byte("packet")...) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + var unpackTime, packTime monotime.Time + done := make(chan struct{}) + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).DoAndReturn( + func(t monotime.Time, bytes []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error) { + unpackTime = monotime.Now() + return protocol.PacketNumber(1), protocol.PacketNumberLen1, protocol.KeyPhaseZero, []byte{0} /* PADDING */, nil + }, + ) + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack) + + switch expectKeepAlive { + case true: + // record the time of the keep-alive is sent + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + packTime = monotime.Now() + close(done) + return shortHeaderPacket{}, errNothingToPack + }, + ) + tc.conn.handlePacket(receivedPacket{data: buf.Data, buffer: buf, rcvTime: monotime.Now(), remoteAddr: tc.remoteAddr}) + select { + case <-done: + // the keep-alive packet should be sent after half the idle timeout + require.Equal(t, unpackTime.Add(idleTimeout/2), packTime) + case <-time.After(idleTimeout): + t.Fatal("timeout") + } + case false: // if keep-alives are disabled, the connection will run into an idle timeout + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.handlePacket(receivedPacket{data: buf.Data, buffer: buf, rcvTime: monotime.Now(), remoteAddr: tc.remoteAddr}) + } + + // test teardown + if expectKeepAlive { + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + } + + synctest.Wait() + + select { + case err := <-errChan: + if expectKeepAlive { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, &IdleTimeoutError{}) + } + case <-time.After(time.Hour): + t.Fatal("timeout") + } + }) +} + +func TestConnectionACKTimer(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + &Config{MaxIdleTimeout: time.Second}, + false, + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + ) + const alarmTimeout = 500 * time.Millisecond + + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).AnyTimes() + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + sph.EXPECT().ECNMode(gomock.Any()).AnyTimes() + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + // Set initial alarm timeout far in the future + _ = tc.receivedPacketHandler().ReceivedPacket(1, protocol.ECNNon, protocol.Encryption1RTT, monotime.Now().Add(time.Hour), true) + + var times []monotime.Time + done := make(chan struct{}, 5) + var calls []any + + for range 2 { + calls = append(calls, tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buf *packetBuffer, _ protocol.ByteCount, _ monotime.Time, _ protocol.Version) (shortHeaderPacket, error) { + buf.Data = append(buf.Data, []byte("foobar")...) + times = append(times, monotime.Now()) + rph := tc.receivedPacketHandler() + if len(times) == 1 { + // After first packet is sent, set alarm timeout for the next iteration + // Get the ACK frame to reset state, then receive a new packet to set alarm + _ = rph.GetAckFrame(protocol.Encryption1RTT, monotime.Now(), false) + alarmRcvTime := monotime.Now().Add(alarmTimeout - protocol.MaxAckDelay) + _ = rph.ReceivedPacket(2, protocol.ECNNon, protocol.Encryption1RTT, alarmRcvTime, true) + } else { + // After second packet is sent, set alarm timeout far in the future + _ = rph.GetAckFrame(protocol.Encryption1RTT, monotime.Now(), false) + _ = rph.ReceivedPacket(3, protocol.ECNNon, protocol.Encryption1RTT, monotime.Now().Add(time.Hour), true) + } + return shortHeaderPacket{Frames: []ackhandler.Frame{{Frame: &wire.PingFrame{}}}, Length: 6}, nil + }, + )) + calls = append(calls, tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(*packetBuffer, protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, error) { + done <- struct{}{} + return shortHeaderPacket{}, errNothingToPack + }, + )) + } + gomock.InOrder(calls...) + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + for range 2 { + synctest.Wait() + + select { + case <-done: + case <-time.After(time.Hour): + t.Fatal("timeout") + } + } + + assert.Len(t, times, 2) + require.Equal(t, times[0].Add(alarmTimeout), times[1]) + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + }) +} + +// Send a GSO batch, until we have no more data to send. +func TestConnectionGSOBatch(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + true, + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + ) + + // allow packets to be sent + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).AnyTimes() + sph.EXPECT().TimeUntilSend().AnyTimes() + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().ECNMode(gomock.Any()).Return(protocol.ECT1).AnyTimes() + + maxPacketSize := tc.conn.maxPacketSize() + var expectedData []byte + for i := range 4 { + data := bytes.Repeat([]byte{byte(i)}, int(maxPacketSize)) + expectedData = append(expectedData, data...) + + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + buffer.Data = append(buffer.Data, data...) + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(i)}, nil + }, + ) + } + done := make(chan struct{}) + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack) + tc.sendConn.EXPECT().Write(expectedData, uint16(maxPacketSize), protocol.ECT1).DoAndReturn( + func([]byte, uint16, protocol.ECN) error { close(done); return nil }, + ) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case <-done: + default: + t.Fatal("should have sent a packet") + } + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + }) +} + +// Send a GSO batch, until a packet smaller than the maximum size is packed +func TestConnectionGSOBatchPacketSize(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + true, + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + ) + + // allow packets to be sent + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).AnyTimes() + sph.EXPECT().TimeUntilSend().AnyTimes() + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().ECNMode(gomock.Any()).Return(protocol.ECT1).AnyTimes() + + maxPacketSize := tc.conn.maxPacketSize() + var expectedData []byte + var calls []any + for i := range 4 { + var data []byte + if i == 3 { + data = bytes.Repeat([]byte{byte(i)}, int(maxPacketSize-1)) + } else { + data = bytes.Repeat([]byte{byte(i)}, int(maxPacketSize)) + } + expectedData = append(expectedData, data...) + + calls = append(calls, tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + buffer.Data = append(buffer.Data, data...) + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(10 + i)}, nil + }, + )) + } + // The smaller (fourth) packet concluded this GSO batch, but the send loop will immediately start composing the next batch. + // We therefore send a "foobar", so we can check that we're actually generating two GSO batches. + calls = append(calls, + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + buffer.Data = append(buffer.Data, []byte("foobar")...) + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(14)}, nil + }, + ), + ) + calls = append(calls, + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack), + ) + gomock.InOrder(calls...) + + done := make(chan struct{}) + gomock.InOrder( + tc.sendConn.EXPECT().Write(expectedData, uint16(maxPacketSize), protocol.ECT1), + tc.sendConn.EXPECT().Write([]byte("foobar"), uint16(maxPacketSize), protocol.ECT1).DoAndReturn( + func([]byte, uint16, protocol.ECN) error { close(done); return nil }, + ), + ) + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case <-done: + default: + t.Fatal("should have sent a packet") + } + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + }) +} + +func TestConnectionGSOBatchECN(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + true, + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + ) + + // allow packets to be sent + ecnMode := protocol.ECT1 + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).AnyTimes() + sph.EXPECT().TimeUntilSend().AnyTimes() + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().ECNMode(gomock.Any()).DoAndReturn(func(bool) protocol.ECN { return ecnMode }).AnyTimes() + + // 3. Send a GSO batch, until the ECN marking changes. + var expectedData []byte + var calls []any + maxPacketSize := tc.conn.maxPacketSize() + for i := range 3 { + data := bytes.Repeat([]byte{byte(i)}, int(maxPacketSize)) + expectedData = append(expectedData, data...) + + calls = append(calls, tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + buffer.Data = append(buffer.Data, data...) + if i == 2 { + ecnMode = protocol.ECNCE + } + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(20 + i)}, nil + }, + )) + } + // The smaller (fourth) packet concluded this GSO batch, but the send loop will immediately start composing the next batch. + // We therefore send a "foobar", so we can check that we're actually generating two GSO batches. + calls = append(calls, + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + buffer.Data = append(buffer.Data, []byte("foobar")...) + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(24)}, nil + }, + ), + ) + calls = append(calls, + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(shortHeaderPacket{}, errNothingToPack), + ) + gomock.InOrder(calls...) + + done3 := make(chan struct{}) + tc.sendConn.EXPECT().Write(expectedData, uint16(maxPacketSize), protocol.ECT1) + tc.sendConn.EXPECT().Write([]byte("foobar"), uint16(maxPacketSize), protocol.ECNCE).DoAndReturn( + func([]byte, uint16, protocol.ECN) error { close(done3); return nil }, + ) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case <-done3: + default: + t.Fatal("should have sent a packet") + } + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + }) +} + +func TestConnectionPTOProbePackets(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testConnectionPTOProbePackets(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testConnectionPTOProbePackets(t, protocol.EncryptionHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testConnectionPTOProbePackets(t, protocol.Encryption1RTT) + }) +} + +func testConnectionPTOProbePackets(t *testing.T, encLevel protocol.EncryptionLevel) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptSentPacketHandler(sph), + ) + + var sendMode ackhandler.SendMode + switch encLevel { + case protocol.EncryptionInitial: + sendMode = ackhandler.SendPTOInitial + case protocol.EncryptionHandshake: + sendMode = ackhandler.SendPTOHandshake + case protocol.Encryption1RTT: + sendMode = ackhandler.SendPTOAppData + } + + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().TimeUntilSend().AnyTimes() + sph.EXPECT().SendMode(gomock.Any()).Return(sendMode) + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendNone) + sph.EXPECT().ECNMode(gomock.Any()) + sph.EXPECT().QueueProbePacket(encLevel).Return(false) + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()) + + tc.packer.EXPECT().PackPTOProbePacket(encLevel, gomock.Any(), true, gomock.Any(), protocol.Version1).DoAndReturn( + func(protocol.EncryptionLevel, protocol.ByteCount, bool, monotime.Time, protocol.Version) (*coalescedPacket, error) { + return &coalescedPacket{ + buffer: getPacketBuffer(), + shortHdrPacket: &shortHeaderPacket{PacketNumber: 1}, + }, nil + }, + ) + done := make(chan struct{}) + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).Do( + func([]byte, uint16, protocol.ECN) error { close(done); return nil }, + ) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have timed out") + } + }) +} + +func TestConnectionCongestionControl(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + false, + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + ) + + sph.EXPECT().TimeUntilSend().AnyTimes() + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().ECNMode(true).AnyTimes() + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).Times(2) + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAck).MaxTimes(1) + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Times(2) + // Since we're already sending out packets, we don't expect any calls to PackAckOnlyPacket + for i := range 2 { + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(buffer *packetBuffer, count protocol.ByteCount, t monotime.Time, version protocol.Version) (shortHeaderPacket, error) { + buffer.Data = append(buffer.Data, []byte("foobar")...) + return shortHeaderPacket{PacketNumber: protocol.PacketNumber(i)}, nil + }, + ) + } + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()) + done1 := make(chan struct{}) + tc.sendConn.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).Do( + func([]byte, uint16, protocol.ECN) error { close(done1); return nil }, + ) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case <-done1: + default: + t.Fatal("should have sent a packet") + } + require.True(t, mockCtrl.Satisfied()) + + // Now that we're congestion limited, we can only send an ack-only packet + done2 := make(chan struct{}) + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAck) + tc.packer.EXPECT().PackAckOnlyPacket(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + close(done2) + return shortHeaderPacket{}, nil, errNothingToPack + }, + ) + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case <-done2: + default: + t.Fatal("should have sent an ack-only packet") + } + require.True(t, mockCtrl.Satisfied()) + + // If the send mode is "none", we can't even send an ack-only packet + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendNone) + tc.conn.scheduleSending() + synctest.Wait() // make sure there are no calls to the packer + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + }) +} + +func TestConnectionSendQueue(t *testing.T) { + t.Run("with GSO", func(t *testing.T) { + testConnectionSendQueue(t, true) + }) + t.Run("without GSO", func(t *testing.T) { + testConnectionSendQueue(t, false) + }) +} + +func testConnectionSendQueue(t *testing.T, enableGSO bool) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + sph := mockackhandler.NewMockSentPacketHandler(mockCtrl) + sender := NewMockSender(mockCtrl) + tc := newServerTestConnection(t, + mockCtrl, + nil, + enableGSO, + connectionOptSender(sender), + connectionOptHandshakeConfirmed(), + connectionOptSentPacketHandler(sph), + ) + + sender.EXPECT().Run().MaxTimes(1) + sender.EXPECT().WouldBlock() + sender.EXPECT().WouldBlock().Return(true).Times(2) + available := make(chan struct{}) + blocked := make(chan struct{}) + sender.EXPECT().Available().DoAndReturn( + func() <-chan struct{} { + close(blocked) + return available + }, + ) + sph.EXPECT().GetLossDetectionTimeout().AnyTimes() + sph.EXPECT().SentPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()) + sph.EXPECT().SendMode(gomock.Any()).Return(ackhandler.SendAny).AnyTimes() + sph.EXPECT().ECNMode(gomock.Any()).AnyTimes() + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return( + shortHeaderPacket{PacketNumber: protocol.PacketNumber(1)}, nil, + ) + sender.EXPECT().Send(gomock.Any(), gomock.Any(), gomock.Any()) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + tc.conn.scheduleSending() + + synctest.Wait() + + select { + case <-blocked: + default: + t.Fatal("should have blocked") + } + require.True(t, mockCtrl.Satisfied()) + + // now make room in the send queue + sender.EXPECT().WouldBlock().AnyTimes() + unblocked := make(chan struct{}) + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(*packetBuffer, protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, error) { + close(unblocked) + return shortHeaderPacket{}, errNothingToPack + }, + ) + available <- struct{}{} + + synctest.Wait() + + select { + case <-unblocked: + default: + t.Fatal("should have unblocked") + } + + // test teardown + sender.EXPECT().Close() + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + }) +} + +func getVersionNegotiationPacket(src, dest protocol.ConnectionID, versions []protocol.Version) receivedPacket { + b := wire.ComposeVersionNegotiation( + protocol.ArbitraryLenConnectionID(src.Bytes()), + protocol.ArbitraryLenConnectionID(dest.Bytes()), + versions, + ) + return receivedPacket{ + rcvTime: monotime.Now(), + data: b, + buffer: getPacketBuffer(), + } +} + +func TestConnectionVersionNegotiation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newClientTestConnection(t, mockCtrl, nil, false, connectionOptTracer(&eventRecorder)) + + tc.packer.EXPECT().PackCoalescedPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes() + tc.connRunner.EXPECT().Remove(gomock.Any()) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + vnp := getVersionNegotiationPacket( + tc.destConnID, + tc.srcConnID, + []protocol.Version{1234, protocol.Version2}, + ) + // the version negotiation packet might contained greased versions + _, _, vnpVersions, err := wire.ParseVersionNegotiationPacket(vnp.data) + require.NoError(t, err) + tc.conn.handlePacket(vnp) + + synctest.Wait() + + select { + case err := <-errChan: + var rerr *errCloseForRecreating + require.ErrorAs(t, err, &rerr) + require.Equal(t, rerr.nextVersion, protocol.Version2) + default: + t.Fatal("should have received a Version Negotiation packet") + } + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionNegotiationReceived{ + Header: qlog.PacketHeaderVersionNegotiation{ + SrcConnectionID: protocol.ArbitraryLenConnectionID(tc.destConnID.Bytes()), + DestConnectionID: protocol.ArbitraryLenConnectionID(tc.srcConnID.Bytes()), + }, + SupportedVersions: vnpVersions, + }, + qlog.VersionInformation{ + ServerVersions: vnpVersions, + ClientVersions: []qlog.Version{protocol.Version1, protocol.Version2}, + ChosenVersion: protocol.Version2, + }, + }, + eventRecorder.Events(qlog.VersionNegotiationReceived{}, qlog.VersionInformation{}), + ) + }) +} + +func TestConnectionVersionNegotiationNoMatch(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newClientTestConnection(t, + mockCtrl, + &Config{Versions: []protocol.Version{protocol.Version1}}, + false, + connectionOptTracer(&eventRecorder), + ) + + tc.packer.EXPECT().PackCoalescedPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes() + tc.connRunner.EXPECT().Remove(gomock.Any()) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + vnp := getVersionNegotiationPacket( + tc.destConnID, + tc.srcConnID, + []protocol.Version{protocol.Version2}, + ) + _, _, vnpVersions, err := wire.ParseVersionNegotiationPacket(vnp.data) + require.NoError(t, err) + tc.conn.handlePacket(vnp) + + synctest.Wait() + + select { + case err := <-errChan: + var verr *VersionNegotiationError + require.ErrorAs(t, err, &verr) + require.Contains(t, verr.Theirs, protocol.Version2) + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionNegotiationReceived{ + Header: qlog.PacketHeaderVersionNegotiation{ + SrcConnectionID: protocol.ArbitraryLenConnectionID(tc.destConnID.Bytes()), + DestConnectionID: protocol.ArbitraryLenConnectionID(tc.srcConnID.Bytes()), + }, + SupportedVersions: vnpVersions, + }, + qlog.ConnectionClosed{ + Initiator: qlog.InitiatorLocal, + Trigger: qlog.ConnectionCloseTriggerVersionMismatch, + }, + }, + eventRecorder.Events(qlog.VersionNegotiationReceived{}, qlog.ConnectionClosed{}), + ) + default: + t.Fatal("should have received a Version Negotiation packet") + } + }) +} + +func TestConnectionVersionNegotiationInvalidPackets(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + tc := newClientTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + ) + + // offers the current version + vnp := getVersionNegotiationPacket( + tc.destConnID, + tc.srcConnID, + []protocol.Version{1234, protocol.Version1}, + ) + wasProcessed, err := tc.conn.handleOnePacket(vnp, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: int(vnp.Size())}, + Trigger: qlog.PacketDropUnexpectedVersion, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + require.True(t, mockCtrl.Satisfied()) + eventRecorder.Clear() + + // unparseable, since it's missing 2 bytes + vnp.data = vnp.data[:len(vnp.data)-2] + wasProcessed, err = tc.conn.handleOnePacket(vnp, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: int(vnp.Size())}, + Trigger: qlog.PacketDropHeaderParseError, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) +} + +func getRetryPacket(t *testing.T, src, dest, origDest protocol.ConnectionID, token []byte) receivedPacket { + hdr := wire.Header{ + Type: protocol.PacketTypeRetry, + SrcConnectionID: src, + DestConnectionID: dest, + Token: token, + Version: protocol.Version1, + } + b, err := (&wire.ExtendedHeader{Header: hdr}).Append(nil, protocol.Version1) + require.NoError(t, err) + tag := handshake.GetRetryIntegrityTag(b, origDest, protocol.Version1) + b = append(b, tag[:]...) + return receivedPacket{ + rcvTime: monotime.Now(), + data: b, + buffer: getPacketBuffer(), + } +} + +func TestConnectionRetryDrops(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + unpacker := NewMockUnpacker(mockCtrl) + tc := newClientTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + connectionOptUnpacker(unpacker), + ) + + newConnID := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}) + + // invalid integrity tag + retry := getRetryPacket(t, newConnID, tc.srcConnID, tc.destConnID, []byte("foobar")) + retry.data[len(retry.data)-1]++ + wasProcessed, err := tc.conn.handleOnePacket(retry, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: newConnID, + DestConnectionID: tc.srcConnID, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(retry.Size())}, + Trigger: qlog.PacketDropPayloadDecryptError, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + // receive a retry that doesn't change the connection ID + retry = getRetryPacket(t, tc.destConnID, tc.srcConnID, tc.destConnID, []byte("foobar")) + wasProcessed, err = tc.conn.handleOnePacket(retry, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: tc.destConnID, + DestConnectionID: tc.srcConnID, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(retry.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) +} + +func TestConnectionRetryAfterReceivedPacket(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + unpacker := NewMockUnpacker(mockCtrl) + tc := newClientTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + connectionOptUnpacker(unpacker), + ) + + // receive a regular packet + regular := getPacketWithPacketType(t, tc.srcConnID, protocol.PacketTypeInitial, 200) + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return( + &unpackedPacket{ + hdr: &wire.ExtendedHeader{Header: wire.Header{Type: protocol.PacketTypeInitial}}, + encryptionLevel: protocol.EncryptionInitial, + }, nil, + ) + wasProcessed, err := tc.conn.handleOnePacket(receivedPacket{ + data: regular, + buffer: getPacketBuffer(), + rcvTime: monotime.Now(), + remoteAddr: tc.remoteAddr, + }, 0) + require.NoError(t, err) + require.True(t, wasProcessed) + + require.Len(t, eventRecorder.Events(qlog.PacketReceived{}), 1) + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionInformation{ + ChosenVersion: protocol.Version1, + ClientVersions: tc.conn.config.Versions, + }, + }, + eventRecorder.Events(qlog.VersionInformation{}), + ) + eventRecorder.Clear() + + // receive a retry + retry := getRetryPacket(t, tc.destConnID, tc.srcConnID, tc.destConnID, []byte("foobar")) + wasProcessed, err = tc.conn.handleOnePacket(retry, 0) + require.NoError(t, err) + require.False(t, wasProcessed) + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: tc.conn.origDestConnID, + DestConnectionID: tc.srcConnID, + Version: tc.conn.version, + }, + Raw: qlog.RawInfo{Length: int(retry.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + eventRecorder.Clear() +} + +func TestConnectionConnectionIDChanges(t *testing.T) { + t.Run("with retry", func(t *testing.T) { + testConnectionConnectionIDChanges(t, true) + }) + t.Run("without retry", func(t *testing.T) { + testConnectionConnectionIDChanges(t, false) + }) +} + +func testConnectionConnectionIDChanges(t *testing.T, sendRetry bool) { + synctest.Test(t, func(t *testing.T) { + makeInitialPacket := func(t *testing.T, hdr *wire.ExtendedHeader) []byte { + t.Helper() + data, err := hdr.Append(nil, protocol.Version1) + require.NoError(t, err) + data = append(data, make([]byte, hdr.Length-protocol.ByteCount(hdr.PacketNumberLen))...) + return data + } + + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + unpacker := NewMockUnpacker(mockCtrl) + tc := newClientTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + connectionOptUnpacker(unpacker), + ) + + dstConnID := tc.destConnID + b := make([]byte, 3*10) + rand.Read(b) + newConnID := protocol.ParseConnectionID(b[:11]) + newConnID2 := protocol.ParseConnectionID(b[11:20]) + + tc.packer.EXPECT().PackCoalescedPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes() + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + require.Equal(t, dstConnID, tc.conn.connIDManager.Get()) + + var retryConnID protocol.ConnectionID + if sendRetry { + retryConnID = protocol.ParseConnectionID(b[20:30]) + tc.packer.EXPECT().SetToken([]byte("foobar")) + + retry := getRetryPacket(t, retryConnID, tc.srcConnID, tc.destConnID, []byte("foobar")) + tc.conn.handlePacket(retry) + + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: retryConnID, + DestConnectionID: dstConnID, + Version: protocol.Version1, + Token: &qlog.Token{Raw: []byte("foobar")}, + }, + Raw: qlog.RawInfo{Length: int(retry.Size())}, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketDropped{}), + ) + } + eventRecorder.Clear() + + // Send the first packet. The server changes the connection ID to newConnID. + hdr1 := wire.ExtendedHeader{ + Header: wire.Header{ + SrcConnectionID: newConnID, + DestConnectionID: tc.srcConnID, + Type: protocol.PacketTypeInitial, + Length: 200, + Version: protocol.Version1, + }, + PacketNumber: 1, + PacketNumberLen: protocol.PacketNumberLen2, + } + hdr2 := hdr1 + hdr2.SrcConnectionID = newConnID2 + + unpacker.EXPECT().UnpackLongHeader(gomock.Any(), gomock.Any()).Return( + &unpackedPacket{hdr: &hdr1, encryptionLevel: protocol.EncryptionInitial}, nil, + ) + eventRecorder.Clear() + packet1 := getLongHeaderPacket(t, tc.remoteAddr, &hdr1, make([]byte, 198)) + tc.conn.handlePacket(packet1) + + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + SrcConnectionID: newConnID, + DestConnectionID: tc.srcConnID, + PacketNumber: 1, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(packet1.Size()), PayloadLength: int(hdr1.Length)}, + DatagramPayloadChecksum: qlog.CalculateDatagramPayloadChecksum(packet1.data), + Frames: []qlog.Frame{}, + }, + }, + eventRecorder.Events(qlog.PacketReceived{}, qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + // Send the second packet. We refuse to accept it, because the connection ID is changed again. + packet2 := receivedPacket{data: makeInitialPacket(t, &hdr2), buffer: getPacketBuffer(), rcvTime: monotime.Now(), remoteAddr: tc.remoteAddr} + tc.conn.handlePacket(packet2) + + synctest.Wait() + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(packet2.Size())}, + DatagramPayloadChecksum: qlog.CalculateDatagramPayloadChecksum(packet2.data), + Trigger: qlog.PacketDropUnknownConnectionID, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}, qlog.PacketReceived{}), + ) + // the connection ID should not have changed + require.Equal(t, newConnID, tc.conn.connIDManager.Get()) + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()) + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have shut down") + } + }) +} + +// When the connection is closed before sending the first packet, +// we don't send a CONNECTION_CLOSE. +// This can happen if there's something wrong the tls.Config, and +// crypto/tls refuses to start the handshake. +func TestConnectionEarlyClose(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + var eventRecorder events.Recorder + cryptoSetup := mocks.NewMockCryptoSetup(mockCtrl) + tc := newClientTestConnection(t, + mockCtrl, + nil, + false, + connectionOptTracer(&eventRecorder), + connectionOptCryptoSetup(cryptoSetup), + ) + + tc.conn.sentFirstPacket = false + cryptoSetup.EXPECT().StartHandshake(gomock.Any()).Do(func(context.Context) error { + tc.conn.closeLocal(errors.New("early error")) + return nil + }) + cryptoSetup.EXPECT().NextEvent().Return(handshake.Event{Kind: handshake.EventNoEvent}) + cryptoSetup.EXPECT().Close() + tc.connRunner.EXPECT().Remove(gomock.Any()) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + synctest.Wait() + + select { + case err := <-errChan: + require.Error(t, err) + require.ErrorContains(t, err, "early error") + code := qerr.InternalError + require.Equal(t, + []qlogwriter.Event{ + qlog.ConnectionClosed{ + Initiator: qlog.InitiatorLocal, + ConnectionError: &code, + Reason: "early error", + }, + }, + eventRecorder.Events(qlog.ConnectionClosed{}), + ) + default: + t.Fatal("should have shut down") + } + }) +} + +func TestConnectionPathValidation(t *testing.T) { + t.Run("NAT rebinding", func(t *testing.T) { + testConnectionPathValidation(t, true) + }) + + t.Run("intentional migration", func(t *testing.T) { + testConnectionPathValidation(t, false) + }) +} + +func testConnectionPathValidation(t *testing.T, isNATRebinding bool) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + unpacker := NewMockUnpacker(mockCtrl) + tc := newServerTestConnection( + t, + mockCtrl, + nil, + false, + connectionOptUnpacker(unpacker), + connectionOptHandshakeConfirmed(), + connectionOptRTT(time.Second), + ) + require.NoError(t, tc.conn.handleTransportParameters(&wire.TransportParameters{MaxUDPPayloadSize: 1456})) + + newRemoteAddr := &net.UDPAddr{IP: net.IPv4(192, 168, 1, 1), Port: 1234} + require.NotEqual(t, tc.remoteAddr, newRemoteAddr) + + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + probeSent := make(chan struct{}) + var pathChallenge *wire.PathChallengeFrame + payload := []byte{0} // PADDING frame + if isNATRebinding { + payload = []byte{1} // PING frame + } + gomock.InOrder( + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return( + protocol.PacketNumber(10), protocol.PacketNumberLen2, protocol.KeyPhaseZero, payload, nil, + ), + tc.packer.EXPECT().PackPathProbePacket(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(_ protocol.ConnectionID, frames []ackhandler.Frame, _ protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + pathChallenge = frames[0].Frame.(*wire.PathChallengeFrame) + return shortHeaderPacket{IsPathProbePacket: true}, getPacketBuffer(), nil + }, + ), + tc.sendConn.EXPECT().WriteTo(gomock.Any(), newRemoteAddr, packetInfo{}).DoAndReturn( + func([]byte, net.Addr, packetInfo) error { close(probeSent); return nil }, + ), + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return( + shortHeaderPacket{}, errNothingToPack, + ), + ) + + tc.conn.handlePacket(receivedPacket{ + data: make([]byte, 10), + buffer: getPacketBuffer(), + remoteAddr: newRemoteAddr, + rcvTime: monotime.Now(), + }) + + synctest.Wait() + + select { + case <-probeSent: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // Receive a packed containing a PATH_RESPONSE frame. + // Only if the first packet received on the path was a probing packet + // (i.e. we're dealing with a NAT rebinding), this makes us switch to the new path. + migrated := make(chan struct{}) + data, err := (&wire.PathResponseFrame{Data: pathChallenge.Data}).Append(nil, protocol.Version1) + require.NoError(t, err) + calls := []any{ + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return( + protocol.PacketNumber(11), protocol.PacketNumberLen2, protocol.KeyPhaseZero, data, nil, + ), + } + if isNATRebinding { + calls = append(calls, + tc.sendConn.EXPECT().ChangeRemoteAddr(newRemoteAddr, gomock.Any()).Do( + func(net.Addr, packetInfo) { close(migrated) }, + ), + ) + } + calls = append(calls, + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return( + shortHeaderPacket{}, errNothingToPack, + ).MaxTimes(1), + ) + gomock.InOrder(calls...) + require.Equal(t, tc.remoteAddr, tc.conn.RemoteAddr()) + // the PATH_RESPONSE can be sent on the old path, if the client is just probing the new path + addr := tc.remoteAddr + if isNATRebinding { + addr = newRemoteAddr + } + tc.conn.handlePacket(receivedPacket{ + data: make([]byte, 100), + buffer: getPacketBuffer(), + remoteAddr: addr, + rcvTime: monotime.Now(), + }) + + synctest.Wait() + + if !isNATRebinding { + // If the first packet was a probing packet, we only switch to the new path when we + // receive a non-probing packet on that path. + select { + case <-migrated: + t.Fatal("didn't expect a migration yet") + default: + } + + payload := []byte{1} // PING frame + payload, err = (&wire.PathResponseFrame{Data: pathChallenge.Data}).Append(payload, protocol.Version1) + require.NoError(t, err) + gomock.InOrder( + unpacker.EXPECT().UnpackShortHeader(gomock.Any(), gomock.Any()).Return( + protocol.PacketNumber(12), protocol.PacketNumberLen2, protocol.KeyPhaseZero, payload, nil, + ), + tc.sendConn.EXPECT().ChangeRemoteAddr(newRemoteAddr, gomock.Any()).Do( + func(net.Addr, packetInfo) { close(migrated) }, + ), + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return( + shortHeaderPacket{}, errNothingToPack, + ).MaxTimes(1), + ) + tc.conn.handlePacket(receivedPacket{ + data: make([]byte, 100), + buffer: getPacketBuffer(), + remoteAddr: newRemoteAddr, + rcvTime: monotime.Now(), + }) + } + + synctest.Wait() + + select { + case <-migrated: + default: + t.Fatal("should have migrated") + } + + // test teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.conn.destroy(nil) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("should have shut down") + } + }) +} + +func TestConnectionMigrationServer(t *testing.T) { + tc := newServerTestConnection(t, nil, nil, false) + _, err := tc.conn.AddPath(&Transport{}) + require.Error(t, err) + require.ErrorContains(t, err, "server cannot initiate connection migration") +} + +func TestConnectionMigration(t *testing.T) { + t.Run("disabled", func(t *testing.T) { + testConnectionMigration(t, false) + }) + + t.Run("enabled", func(t *testing.T) { + testConnectionMigration(t, true) + }) +} + +func testConnectionMigration(t *testing.T, enabled bool) { + tc := newClientTestConnection(t, nil, nil, false, connectionOptHandshakeConfirmed()) + require.NoError(t, tc.conn.handleTransportParameters(&wire.TransportParameters{ + InitialSourceConnectionID: tc.destConnID, + OriginalDestinationConnectionID: tc.destConnID, + DisableActiveMigration: !enabled, + })) + + tr := &Transport{ + Conn: newUDPConnLocalhost(t), + StatelessResetKey: &StatelessResetKey{}, + } + defer tr.Close() + path, err := tc.conn.AddPath(tr) + if !enabled { + require.Error(t, err) + require.ErrorContains(t, err, "server disabled connection migration") + return + } + require.NoError(t, err) + require.NotNil(t, path) + + tc.packer.EXPECT().AppendPacket(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return( + shortHeaderPacket{}, errNothingToPack, + ).AnyTimes() + packedProbe := make(chan struct{}) + tc.packer.EXPECT().PackPathProbePacket(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(protocol.ConnectionID, []ackhandler.Frame, protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + defer close(packedProbe) + return shortHeaderPacket{IsPathProbePacket: true}, getPacketBuffer(), nil + }, + ).AnyTimes() + tc.connRunner.EXPECT().AddResetToken(gomock.Any(), gomock.Any()) + // add a new connection ID, so the path can be probed + _, err = tc.conn.handleFrame(&wire.NewConnectionIDFrame{ + SequenceNumber: 1, + ConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + }, protocol.EncryptionInitial, tc.destConnID, monotime.Now()) + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { errChan <- tc.conn.run() }() + + // Adding the path initialized the transport. + // We can test this by triggering a stateless reset. + conn := newUDPConnLocalhost(t) + _, err = conn.WriteTo(append([]byte{0x40}, make([]byte, 100)...), tr.Conn.LocalAddr()) + require.NoError(t, err) + conn.SetReadDeadline(time.Now().Add(time.Second)) + _, _, err = conn.ReadFrom(make([]byte, 100)) + require.NoError(t, err) + + go func() { path.Probe(context.Background()) }() + select { + case <-packedProbe: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // teardown + tc.connRunner.EXPECT().Remove(gomock.Any()).AnyTimes() + tc.connRunner.EXPECT().RemoveResetToken(gomock.Any()).MaxTimes(1) + tc.conn.destroy(nil) + select { + case <-errChan: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnectionDatagrams(t *testing.T) { + t.Run("disabled", func(t *testing.T) { + testConnectionDatagrams(t, false) + }) + t.Run("enabled", func(t *testing.T) { + testConnectionDatagrams(t, true) + }) +} + +func testConnectionDatagrams(t *testing.T, enabled bool) { + tc := newServerTestConnection(t, nil, &Config{EnableDatagrams: enabled}, false) + + data, err := (&wire.DatagramFrame{Data: []byte("foo"), DataLenPresent: true}).Append(nil, protocol.Version1) + require.NoError(t, err) + data, err = (&wire.DatagramFrame{Data: []byte("bar")}).Append(data, protocol.Version1) + require.NoError(t, err) + _, _, _, err = tc.conn.handleFrames(data, protocol.ConnectionID{}, protocol.Encryption1RTT, nil, monotime.Now()) + + if !enabled { + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.FrameEncodingError, FrameType: uint64(wire.FrameTypeDatagramWithLength)}) + return + } + + require.NoError(t, err) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + d, err := tc.conn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("foo"), d) + d, err = tc.conn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("bar"), d) +} diff --git a/third_party/quic-go/crypto_stream.go b/third_party/quic-go/crypto_stream.go new file mode 100644 index 0000000..3e6a5bd --- /dev/null +++ b/third_party/quic-go/crypto_stream.go @@ -0,0 +1,283 @@ +package quic + +import ( + "errors" + "fmt" + "io" + "os" + "slices" + "strconv" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +const disableClientHelloScramblingEnv = "QUIC_GO_DISABLE_CLIENTHELLO_SCRAMBLING" + +// The baseCryptoStream is used by the cryptoStream and the initialCryptoStream. +// This allows us to implement different logic for PopCryptoFrame for the two streams. +type baseCryptoStream struct { + queue frameSorter + + highestOffset protocol.ByteCount + finished bool + + writeOffset protocol.ByteCount + writeBuf []byte +} + +func newCryptoStream() *cryptoStream { + return &cryptoStream{baseCryptoStream{queue: *newFrameSorter()}} +} + +func (s *baseCryptoStream) HandleCryptoFrame(f *wire.CryptoFrame) error { + highestOffset := f.Offset + protocol.ByteCount(len(f.Data)) + if maxOffset := highestOffset; maxOffset > protocol.MaxCryptoStreamOffset { + return &qerr.TransportError{ + ErrorCode: qerr.CryptoBufferExceeded, + ErrorMessage: fmt.Sprintf("received invalid offset %d on crypto stream, maximum allowed %d", maxOffset, protocol.MaxCryptoStreamOffset), + } + } + if s.finished { + if highestOffset > s.highestOffset { + // reject crypto data received after this stream was already finished + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "received crypto data after change of encryption level", + } + } + // ignore data with a smaller offset than the highest received + // could e.g. be a retransmission + return nil + } + s.highestOffset = max(s.highestOffset, highestOffset) + return s.queue.Push(f.Data, f.Offset, nil) +} + +// GetCryptoData retrieves data that was received in CRYPTO frames +func (s *baseCryptoStream) GetCryptoData() []byte { + _, data, _ := s.queue.Pop() + return data +} + +func (s *baseCryptoStream) Finish() error { + if s.queue.HasMoreData() { + return &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "encryption level changed, but crypto stream has more data to read", + } + } + s.finished = true + return nil +} + +// Writes writes data that should be sent out in CRYPTO frames +func (s *baseCryptoStream) Write(p []byte) (int, error) { + s.writeBuf = append(s.writeBuf, p...) + return len(p), nil +} + +func (s *baseCryptoStream) HasData() bool { + return len(s.writeBuf) > 0 +} + +// PendingLen returns the number of bytes still queued for sending, which the +// packet packer uses to size the crypto budget of each Initial. +func (s *baseCryptoStream) PendingLen() protocol.ByteCount { + return protocol.ByteCount(len(s.writeBuf)) +} + +// WriteOffset returns the crypto stream offset the next CRYPTO frame will carry. +// Sizing that frame to an exact byte count requires knowing the offset, since it +// determines how many bytes the frame header occupies. +func (s *baseCryptoStream) WriteOffset() protocol.ByteCount { + return s.writeOffset +} + +func (s *baseCryptoStream) PopCryptoFrame(maxLen protocol.ByteCount) *wire.CryptoFrame { + f := &wire.CryptoFrame{Offset: s.writeOffset} + n := min(f.MaxDataLen(maxLen), protocol.ByteCount(len(s.writeBuf))) + if n <= 0 { + return nil + } + f.Data = s.writeBuf[:n] + s.writeBuf = s.writeBuf[n:] + s.writeOffset += n + return f +} + +// PopCryptoFrameTail takes dataLen bytes from the end of what is queued rather +// than the start, so a packet can carry the last part of the ClientHello while +// the middle is still pending. The write offset is unchanged, so later calls to +// PopCryptoFrame continue from the front. +func (s *baseCryptoStream) PopCryptoFrameTail(dataLen protocol.ByteCount) *wire.CryptoFrame { + if dataLen <= 0 || dataLen > protocol.ByteCount(len(s.writeBuf)) { + return nil + } + n := protocol.ByteCount(len(s.writeBuf)) - dataLen + f := &wire.CryptoFrame{ + Offset: s.writeOffset + n, + Data: s.writeBuf[n:], + } + s.writeBuf = s.writeBuf[:n] + return f +} + +type cryptoStream struct { + baseCryptoStream +} + +type clientHelloCut struct { + start protocol.ByteCount + end protocol.ByteCount +} + +type initialCryptoStream struct { + baseCryptoStream + + scramble bool + end protocol.ByteCount + cuts [2]clientHelloCut +} + +// chromeParrot disables quic-go's own ClientHello scrambling. That scheme cuts at +// semantically chosen offsets, which is recognizably quic-go; the packet packer +// applies chaos protection instead, and running both would produce a hybrid +// matching neither. +func newInitialCryptoStream(isClient, chromeParrot bool) *initialCryptoStream { + var scramble bool + if isClient && !chromeParrot { + disabled, err := strconv.ParseBool(os.Getenv(disableClientHelloScramblingEnv)) + scramble = err != nil || !disabled + } + s := &initialCryptoStream{ + baseCryptoStream: baseCryptoStream{queue: *newFrameSorter()}, + scramble: scramble, + } + for i := range len(s.cuts) { + s.cuts[i].start = protocol.InvalidByteCount + s.cuts[i].end = protocol.InvalidByteCount + } + return s +} + +func (s *initialCryptoStream) HasData() bool { + // The ClientHello might be written in multiple parts. + // In order to correctly split the ClientHello, we need the entire ClientHello has been queued. + if s.scramble && s.writeOffset == 0 && s.cuts[0].start == protocol.InvalidByteCount { + return false + } + return s.baseCryptoStream.HasData() +} + +func (s *initialCryptoStream) Write(p []byte) (int, error) { + s.writeBuf = append(s.writeBuf, p...) + if !s.scramble { + return len(p), nil + } + if s.cuts[0].start == protocol.InvalidByteCount { + sniPos, sniLen, echPos, err := findSNIAndECH(s.writeBuf) + if errors.Is(err, io.ErrUnexpectedEOF) { + return len(p), nil + } + if err != nil { + return len(p), err + } + if sniPos == -1 && echPos == -1 { + // Neither SNI nor ECH found. + // There's nothing to scramble. + s.scramble = false + return len(p), nil + } + s.end = protocol.ByteCount(len(s.writeBuf)) + s.cuts[0].start = protocol.ByteCount(sniPos + sniLen/2) // right in the middle + s.cuts[0].end = protocol.ByteCount(sniPos + sniLen) + if echPos > 0 { + // ECH extension found, cut the ECH extension type value (a uint16) in half + start := protocol.ByteCount(echPos + 1) + s.cuts[1].start = start + // cut somewhere (16 bytes), most likely in the ECH extension value + s.cuts[1].end = min(start+16, s.end) + } + slices.SortFunc(s.cuts[:], func(a, b clientHelloCut) int { + if a.start == protocol.InvalidByteCount { + return 1 + } + if a.start > b.start { + return 1 + } + return -1 + }) + } + return len(p), nil +} + +func (s *initialCryptoStream) PopCryptoFrame(maxLen protocol.ByteCount) *wire.CryptoFrame { + if !s.scramble { + return s.baseCryptoStream.PopCryptoFrame(maxLen) + } + + // send out the skipped parts + if s.writeOffset == s.end { + var foundCuts bool + var f *wire.CryptoFrame + for i, c := range s.cuts { + if c.start == protocol.InvalidByteCount { + continue + } + foundCuts = true + if f != nil { + break + } + f = &wire.CryptoFrame{Offset: c.start} + n := min(f.MaxDataLen(maxLen), c.end-c.start) + if n <= 0 { + return nil + } + f.Data = s.writeBuf[c.start : c.start+n] + s.cuts[i].start += n + if s.cuts[i].start == c.end { + s.cuts[i].start = protocol.InvalidByteCount + s.cuts[i].end = protocol.InvalidByteCount + foundCuts = false + } + } + if !foundCuts { + // no more cuts found, we're done sending out everything up until s.end + s.writeBuf = s.writeBuf[s.end:] + s.end = protocol.InvalidByteCount + s.scramble = false + } + return f + } + + nextCut := clientHelloCut{start: protocol.InvalidByteCount, end: protocol.InvalidByteCount} + for _, c := range s.cuts { + if c.start == protocol.InvalidByteCount { + continue + } + if c.start > s.writeOffset { + nextCut = c + break + } + } + f := &wire.CryptoFrame{Offset: s.writeOffset} + maxOffset := nextCut.start + if maxOffset == protocol.InvalidByteCount { + maxOffset = s.end + } + n := min(f.MaxDataLen(maxLen), maxOffset-s.writeOffset) + if n <= 0 { + return nil + } + f.Data = s.writeBuf[s.writeOffset : s.writeOffset+n] + // Don't reslice the writeBuf yet. + // This is done once all parts have been sent out. + s.writeOffset += n + if s.writeOffset == nextCut.start { + s.writeOffset = nextCut.end + } + + return f +} diff --git a/third_party/quic-go/crypto_stream_manager.go b/third_party/quic-go/crypto_stream_manager.go new file mode 100644 index 0000000..a0ff9eb --- /dev/null +++ b/third_party/quic-go/crypto_stream_manager.go @@ -0,0 +1,73 @@ +package quic + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +type cryptoStreamManager struct { + initialStream *initialCryptoStream + handshakeStream *cryptoStream + oneRTTStream *cryptoStream +} + +func newCryptoStreamManager( + initialStream *initialCryptoStream, + handshakeStream *cryptoStream, + oneRTTStream *cryptoStream, +) *cryptoStreamManager { + return &cryptoStreamManager{ + initialStream: initialStream, + handshakeStream: handshakeStream, + oneRTTStream: oneRTTStream, + } +} + +func (m *cryptoStreamManager) HandleCryptoFrame(frame *wire.CryptoFrame, encLevel protocol.EncryptionLevel) error { + //nolint:exhaustive // CRYPTO frames cannot be sent in 0-RTT packets. + switch encLevel { + case protocol.EncryptionInitial: + return m.initialStream.HandleCryptoFrame(frame) + case protocol.EncryptionHandshake: + return m.handshakeStream.HandleCryptoFrame(frame) + case protocol.Encryption1RTT: + return m.oneRTTStream.HandleCryptoFrame(frame) + default: + return fmt.Errorf("received CRYPTO frame with unexpected encryption level: %s", encLevel) + } +} + +func (m *cryptoStreamManager) GetCryptoData(encLevel protocol.EncryptionLevel) []byte { + //nolint:exhaustive // CRYPTO frames cannot be sent in 0-RTT packets. + switch encLevel { + case protocol.EncryptionInitial: + return m.initialStream.GetCryptoData() + case protocol.EncryptionHandshake: + return m.handshakeStream.GetCryptoData() + case protocol.Encryption1RTT: + return m.oneRTTStream.GetCryptoData() + default: + panic(fmt.Sprintf("received CRYPTO frame with unexpected encryption level: %s", encLevel)) + } +} + +func (m *cryptoStreamManager) GetPostHandshakeData(maxSize protocol.ByteCount) *wire.CryptoFrame { + if !m.oneRTTStream.HasData() { + return nil + } + return m.oneRTTStream.PopCryptoFrame(maxSize) +} + +func (m *cryptoStreamManager) Drop(encLevel protocol.EncryptionLevel) error { + //nolint:exhaustive // 1-RTT keys should never get dropped. + switch encLevel { + case protocol.EncryptionInitial: + return m.initialStream.Finish() + case protocol.EncryptionHandshake: + return m.handshakeStream.Finish() + default: + panic(fmt.Sprintf("dropped unexpected encryption level: %s", encLevel)) + } +} diff --git a/third_party/quic-go/crypto_stream_manager_test.go b/third_party/quic-go/crypto_stream_manager_test.go new file mode 100644 index 0000000..c4f3392 --- /dev/null +++ b/third_party/quic-go/crypto_stream_manager_test.go @@ -0,0 +1,87 @@ +package quic + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestCryptoStreamManager(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testCryptoStreamManager(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testCryptoStreamManager(t, protocol.EncryptionHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testCryptoStreamManager(t, protocol.Encryption1RTT) + }) +} + +func testCryptoStreamManager(t *testing.T, encLevel protocol.EncryptionLevel) { + initialStream := newInitialCryptoStream(true, false) + handshakeStream := newCryptoStream() + oneRTTStream := newCryptoStream() + csm := newCryptoStreamManager(initialStream, handshakeStream, oneRTTStream) + + require.NoError(t, csm.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("foo")}, encLevel)) + require.NoError(t, csm.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("bar"), Offset: 3}, encLevel)) + var data []byte + for { + b := csm.GetCryptoData(encLevel) + if len(b) == 0 { + break + } + data = append(data, b...) + } + require.Equal(t, []byte("foobar"), data) +} + +func TestCryptoStreamManagerInvalidEncryptionLevel(t *testing.T) { + csm := newCryptoStreamManager(nil, nil, nil) + require.ErrorContains(t, + csm.HandleCryptoFrame(&wire.CryptoFrame{}, protocol.Encryption0RTT), + "received CRYPTO frame with unexpected encryption level", + ) +} + +func TestCryptoStreamManagerDropEncryptionLevel(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testCryptoStreamManagerDropEncryptionLevel(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testCryptoStreamManagerDropEncryptionLevel(t, protocol.EncryptionHandshake) + }) +} + +func testCryptoStreamManagerDropEncryptionLevel(t *testing.T, encLevel protocol.EncryptionLevel) { + initialStream := newInitialCryptoStream(true, false) + handshakeStream := newCryptoStream() + oneRTTStream := newCryptoStream() + csm := newCryptoStreamManager(initialStream, handshakeStream, oneRTTStream) + + require.NoError(t, csm.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("foo")}, encLevel)) + require.ErrorContains(t, csm.Drop(encLevel), "encryption level changed, but crypto stream has more data to read") + + require.Equal(t, []byte("foo"), csm.GetCryptoData(encLevel)) + require.NoError(t, csm.Drop(encLevel)) +} + +func TestCryptoStreamManagerPostHandshake(t *testing.T) { + initialStream := newInitialCryptoStream(true, false) + handshakeStream := newCryptoStream() + oneRTTStream := newCryptoStream() + csm := newCryptoStreamManager(initialStream, handshakeStream, oneRTTStream) + + _, err := oneRTTStream.Write([]byte("foo")) + require.NoError(t, err) + _, err = oneRTTStream.Write([]byte("bar")) + require.NoError(t, err) + require.Equal(t, + &wire.CryptoFrame{Data: []byte("foobar")}, + csm.GetPostHandshakeData(protocol.ByteCount(10)), + ) +} diff --git a/third_party/quic-go/crypto_stream_test.go b/third_party/quic-go/crypto_stream_test.go new file mode 100644 index 0000000..3bfd32d --- /dev/null +++ b/third_party/quic-go/crypto_stream_test.go @@ -0,0 +1,288 @@ +package quic + +import ( + "fmt" + mrand "math/rand/v2" + "os" + "slices" + "strconv" + "strings" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestCryptoStreamDataAssembly(t *testing.T) { + str := newCryptoStream() + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("bar"), Offset: 3})) + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("foo")})) + // receive a retransmission + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("bar"), Offset: 3})) + + var data []byte + for { + b := str.GetCryptoData() + if b == nil { + break + } + data = append(data, b...) + } + require.Equal(t, []byte("foobar"), data) +} + +func TestCryptoStreamMaxOffset(t *testing.T) { + str := newCryptoStream() + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{ + Offset: protocol.MaxCryptoStreamOffset - 5, + Data: []byte("foo"), + })) + require.ErrorIs(t, + str.HandleCryptoFrame(&wire.CryptoFrame{ + Offset: protocol.MaxCryptoStreamOffset - 2, + Data: []byte("bar"), + }), + &qerr.TransportError{ErrorCode: qerr.CryptoBufferExceeded}, + ) +} + +func TestCryptoStreamFinishWithQueuedData(t *testing.T) { + t.Run("with data at current offset", func(t *testing.T) { + str := newCryptoStream() + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("foo")})) + require.Equal(t, []byte("foo"), str.GetCryptoData()) + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("bar"), Offset: 3})) + require.ErrorIs(t, str.Finish(), &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + }) + + t.Run("with data at a higher offset", func(t *testing.T) { + str := newCryptoStream() + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("foobar"), Offset: 20})) + require.ErrorIs(t, str.Finish(), &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + }) +} + +func TestCryptoStreamReceiveDataAfterFinish(t *testing.T) { + str := newCryptoStream() + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("foobar")})) + require.Equal(t, []byte("foobar"), str.GetCryptoData()) + require.NoError(t, str.Finish()) + // receiving a retransmission is ok + require.NoError(t, str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("bar"), Offset: 3})) + // but receiving new data is not + require.ErrorIs(t, + str.HandleCryptoFrame(&wire.CryptoFrame{Data: []byte("baz"), Offset: 4}), + &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}, + ) +} + +func expectedCryptoFrameLen(offset protocol.ByteCount) protocol.ByteCount { + f := &wire.CryptoFrame{Offset: offset} + return f.Length(protocol.Version1) +} + +func TestCryptoStreamWrite(t *testing.T) { + str := newCryptoStream() + + require.False(t, str.HasData()) + _, err := str.Write([]byte("foo")) + require.NoError(t, err) + require.True(t, str.HasData()) + _, err = str.Write([]byte("bar")) + require.NoError(t, err) + _, err = str.Write([]byte("baz")) + require.NoError(t, err) + require.True(t, str.HasData()) + + for i := range expectedCryptoFrameLen(0) { + require.Nil(t, str.PopCryptoFrame(i)) + } + + f := str.PopCryptoFrame(expectedCryptoFrameLen(0) + 1) + require.Equal(t, &wire.CryptoFrame{Data: []byte("f")}, f) + require.True(t, str.HasData()) + f = str.PopCryptoFrame(expectedCryptoFrameLen(1) + 3) + // the three write calls were coalesced into a single frame + require.Equal(t, &wire.CryptoFrame{Offset: 1, Data: []byte("oob")}, f) + f = str.PopCryptoFrame(protocol.MaxByteCount) + require.Equal(t, &wire.CryptoFrame{Offset: 4, Data: []byte("arbaz")}, f) + require.False(t, str.HasData()) +} + +func TestInitialCryptoStreamServer(t *testing.T) { + str := newInitialCryptoStream(false, false) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + + f := str.PopCryptoFrame(expectedCryptoFrameLen(0) + 3) + require.Equal(t, &wire.CryptoFrame{Offset: 0, Data: []byte("foo")}, f) + require.True(t, str.HasData()) + + // append another CRYPTO frame to the existing slice + f = str.PopCryptoFrame(expectedCryptoFrameLen(3) + 3) + require.Equal(t, &wire.CryptoFrame{Offset: 3, Data: []byte("bar")}, f) + require.False(t, str.HasData()) +} + +func reassembleCryptoData(t *testing.T, segments map[protocol.ByteCount][]byte) []byte { + t.Helper() + + var reassembled []byte + var offset protocol.ByteCount + for len(segments) > 0 { + b, ok := segments[offset] + if !ok { + break + } + reassembled = append(reassembled, b...) + delete(segments, offset) + offset = protocol.ByteCount(len(reassembled)) + } + require.Empty(t, segments) + return reassembled +} + +func skipIfDisableScramblingEnvSet(t *testing.T) { + t.Helper() + disabled, err := strconv.ParseBool(os.Getenv(disableClientHelloScramblingEnv)) + if err == nil && disabled { + t.Skip("ClientHello scrambling disabled via " + disableClientHelloScramblingEnv) + } +} + +func TestInitialCryptoStreamClientStatic(t *testing.T) { + skipIfDisableScramblingEnvSet(t) + + str := newInitialCryptoStream(true, false) + clientHello, err := getClientHello("quic-go.net") + require.NoError(t, err) + _, err = str.Write(clientHello) + require.NoError(t, err) + require.True(t, str.HasData()) + _, err = str.Write([]byte("foobar")) + require.NoError(t, err) + + segments := make(map[protocol.ByteCount][]byte) + + f1 := str.PopCryptoFrame(protocol.MaxByteCount) + require.NotNil(t, f1) + segments[f1.Offset] = f1.Data + require.True(t, str.HasData()) + + f2 := str.PopCryptoFrame(protocol.MaxByteCount) + require.NotNil(t, f2) + require.NotContains(t, segments, f2.Offset) + segments[f2.Offset] = f2.Data + require.True(t, str.HasData()) + require.NotEqual(t, f2.Offset, protocol.ByteCount(len(f1.Data))) + + f3 := str.PopCryptoFrame(protocol.MaxByteCount) + require.NotNil(t, f2) + require.NotContains(t, segments, f3.Offset) + segments[f3.Offset] = f3.Data + require.True(t, str.HasData()) + require.NotEqual(t, f3.Offset, protocol.ByteCount(len(f2.Data))) + + f4 := str.PopCryptoFrame(protocol.MaxByteCount) + require.NotNil(t, f4) + require.NotContains(t, segments, f4.Offset) + segments[f4.Offset] = f4.Data + require.Equal(t, []byte("foobar"), f4.Data) + require.False(t, str.HasData()) + require.NotEqual(t, f4.Offset, protocol.ByteCount(len(f3.Data))) + + reassembled := reassembleCryptoData(t, segments) + require.Equal(t, append(clientHello, []byte("foobar")...), reassembled) +} + +func randomDomainName(length int) string { + const alphabet = "abcdefghijklmnopqrstuvwxyz" + b := make([]byte, length) + for i := range b { + if i > 0 && i < length-1 && mrand.IntN(5) == 0 && b[i-1] != '.' { + b[i] = '.' + } else { + b[i] = alphabet[mrand.IntN(len(alphabet))] + } + } + return string(b) +} + +func TestInitialCryptoStreamClientRandomizedSizes(t *testing.T) { + skipIfDisableScramblingEnvSet(t) + + for i := range 100 { + t.Run(fmt.Sprintf("run %d", i), func(t *testing.T) { + var serverName string + if mrand.Int()%4 > 0 { + serverName = randomDomainName(6 + mrand.IntN(20)) + } + var clientHello []byte + if serverName == "" || !strings.Contains(serverName, ".") || mrand.Int()%2 == 0 { + t.Logf("using a ClientHello without ECH, hostname: %q", serverName) + var err error + clientHello, err = getClientHello(serverName) + require.NoError(t, err) + } else { + t.Logf("using a ClientHello with ECH, hostname: %q", serverName) + var err error + clientHello, err = getClientHelloWithECH(serverName) + require.NoError(t, err) + } + testInitialCryptoStreamClientRandomizedSizes(t, clientHello, serverName) + }) + } +} + +func testInitialCryptoStreamClientRandomizedSizes(t *testing.T, clientHello []byte, expectedServerName string) { + str := newInitialCryptoStream(true, false) + + b := slices.Clone(clientHello) + for len(b) > 0 { + n := min(len(b), mrand.IntN(2*len(b))) + _, err := str.Write(b[:n]) + require.NoError(t, err) + b = b[n:] + } + + require.True(t, str.HasData()) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + + segments := make(map[protocol.ByteCount][]byte) + + var frames []*wire.CryptoFrame + for str.HasData() { + // fmt.Println("popping a frame") + var maxSize protocol.ByteCount + if mrand.Int()%4 == 0 { + maxSize = protocol.ByteCount(mrand.IntN(512) + 1) + } else { + maxSize = protocol.ByteCount(mrand.IntN(32) + 1) + } + f := str.PopCryptoFrame(maxSize) + if f == nil { + continue + } + frames = append(frames, f) + require.LessOrEqual(t, f.Length(protocol.Version1), maxSize) + } + t.Logf("received %d frames", len(frames)) + + for _, f := range frames { + t.Logf("offset %d: %d bytes", f.Offset, len(f.Data)) + if expectedServerName != "" { + require.NotContainsf(t, string(f.Data), expectedServerName, "frame at offset %d contains the server name", f.Offset) + } + segments[f.Offset] = f.Data + } + + reassembled := reassembleCryptoData(t, segments) + require.Equal(t, append(clientHello, []byte("foobar")...), reassembled) + if expectedServerName != "" { + require.Contains(t, string(reassembled), expectedServerName) + } +} diff --git a/third_party/quic-go/datagram_queue.go b/third_party/quic-go/datagram_queue.go new file mode 100644 index 0000000..16b2d5d --- /dev/null +++ b/third_party/quic-go/datagram_queue.go @@ -0,0 +1,137 @@ +package quic + +import ( + "context" + "sync" + + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/utils/ringbuffer" + "github.com/apernet/quic-go/internal/wire" +) + +const ( + maxDatagramSendQueueLen = 32 + maxDatagramRcvQueueLen = 128 +) + +type datagramQueue struct { + sendMx sync.Mutex + sendQueue ringbuffer.RingBuffer[*wire.DatagramFrame] + sent chan struct{} // used to notify Add that a datagram was dequeued + + rcvMx sync.Mutex + rcvQueue [][]byte + rcvd chan struct{} // used to notify Receive that a new datagram was received + + closeErr error + closed chan struct{} + + hasData func() + + logger utils.Logger +} + +func newDatagramQueue(hasData func(), logger utils.Logger) *datagramQueue { + return &datagramQueue{ + hasData: hasData, + rcvd: make(chan struct{}, 1), + sent: make(chan struct{}, 1), + closed: make(chan struct{}), + logger: logger, + } +} + +// Add queues a new DATAGRAM frame for sending. +// Up to 32 DATAGRAM frames will be queued. +// Once that limit is reached, Add blocks until the queue size has reduced. +func (h *datagramQueue) Add(f *wire.DatagramFrame) error { + h.sendMx.Lock() + + for { + if h.sendQueue.Len() < maxDatagramSendQueueLen { + h.sendQueue.PushBack(f) + h.sendMx.Unlock() + h.hasData() + return nil + } + select { + case <-h.sent: // drain the queue so we don't loop immediately + default: + } + h.sendMx.Unlock() + select { + case <-h.closed: + return h.closeErr + case <-h.sent: + } + h.sendMx.Lock() + } +} + +// Peek gets the next DATAGRAM frame for sending. +// If actually sent out, Pop needs to be called before the next call to Peek. +func (h *datagramQueue) Peek() *wire.DatagramFrame { + h.sendMx.Lock() + defer h.sendMx.Unlock() + if h.sendQueue.Empty() { + return nil + } + return h.sendQueue.PeekFront() +} + +func (h *datagramQueue) Pop() { + h.sendMx.Lock() + defer h.sendMx.Unlock() + _ = h.sendQueue.PopFront() + select { + case h.sent <- struct{}{}: + default: + } +} + +// HandleDatagramFrame handles a received DATAGRAM frame. +func (h *datagramQueue) HandleDatagramFrame(f *wire.DatagramFrame) { + data := make([]byte, len(f.Data)) + copy(data, f.Data) + var queued bool + h.rcvMx.Lock() + if len(h.rcvQueue) < maxDatagramRcvQueueLen { + h.rcvQueue = append(h.rcvQueue, data) + queued = true + select { + case h.rcvd <- struct{}{}: + default: + } + } + h.rcvMx.Unlock() + if !queued && h.logger.Debug() { + h.logger.Debugf("Discarding received DATAGRAM frame (%d bytes payload)", len(f.Data)) + } +} + +// Receive gets a received DATAGRAM frame. +func (h *datagramQueue) Receive(ctx context.Context) ([]byte, error) { + for { + h.rcvMx.Lock() + if len(h.rcvQueue) > 0 { + data := h.rcvQueue[0] + h.rcvQueue = h.rcvQueue[1:] + h.rcvMx.Unlock() + return data, nil + } + h.rcvMx.Unlock() + select { + case <-h.rcvd: + continue + case <-h.closed: + return nil, h.closeErr + case <-ctx.Done(): + return nil, ctx.Err() + } + } +} + +func (h *datagramQueue) CloseWithError(e error) { + h.closeErr = e + close(h.closed) +} diff --git a/third_party/quic-go/datagram_queue_test.go b/third_party/quic-go/datagram_queue_test.go new file mode 100644 index 0000000..1fd2fac --- /dev/null +++ b/third_party/quic-go/datagram_queue_test.go @@ -0,0 +1,180 @@ +package quic + +import ( + "context" + "testing" + "testing/synctest" + + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDatagramQueuePeekAndPop(t *testing.T) { + var queued []struct{} + queue := newDatagramQueue(func() { queued = append(queued, struct{}{}) }, utils.DefaultLogger) + require.Nil(t, queue.Peek()) + require.Empty(t, queued) + require.NoError(t, queue.Add(&wire.DatagramFrame{Data: []byte("foo")})) + require.Len(t, queued, 1) + require.Equal(t, &wire.DatagramFrame{Data: []byte("foo")}, queue.Peek()) + // calling peek again returns the same datagram + require.Equal(t, &wire.DatagramFrame{Data: []byte("foo")}, queue.Peek()) + queue.Pop() + require.Nil(t, queue.Peek()) +} + +func TestDatagramQueueSendQueueLength(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + queue := newDatagramQueue(func() {}, utils.DefaultLogger) + + for range maxDatagramSendQueueLen { + require.NoError(t, queue.Add(&wire.DatagramFrame{Data: []byte{0}})) + } + errChan := make(chan error, 1) + go func() { errChan <- queue.Add(&wire.DatagramFrame{Data: []byte("foobar")}) }() + + synctest.Wait() + + select { + case <-errChan: + t.Fatal("expected to not receive error") + default: + } + + // peeking doesn't remove the datagram from the queue... + require.NotNil(t, queue.Peek()) + synctest.Wait() + select { + case <-errChan: + t.Fatal("expected to not receive error") + default: + } + + // ...but popping does + queue.Pop() + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + // pop all the remaining datagrams + for range maxDatagramSendQueueLen - 1 { + queue.Pop() + } + f := queue.Peek() + require.NotNil(t, f) + require.Equal(t, &wire.DatagramFrame{Data: []byte("foobar")}, f) + }) +} + +func TestDatagramQueueReceive(t *testing.T) { + queue := newDatagramQueue(func() {}, utils.DefaultLogger) + + // receive frames that were received earlier + queue.HandleDatagramFrame(&wire.DatagramFrame{Data: []byte("foo")}) + queue.HandleDatagramFrame(&wire.DatagramFrame{Data: []byte("bar")}) + data, err := queue.Receive(context.Background()) + require.NoError(t, err) + require.Equal(t, []byte("foo"), data) + data, err = queue.Receive(context.Background()) + require.NoError(t, err) + require.Equal(t, []byte("bar"), data) +} + +func TestDatagramQueueReceiveBlocking(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + queue := newDatagramQueue(func() {}, utils.DefaultLogger) + + // block until a new frame is received + type result struct { + data []byte + err error + } + resultChan := make(chan result, 1) + go func() { + data, err := queue.Receive(context.Background()) + resultChan <- result{data, err} + }() + + synctest.Wait() + + select { + case <-resultChan: + t.Fatal("expected to not receive result") + default: + } + queue.HandleDatagramFrame(&wire.DatagramFrame{Data: []byte("foobar")}) + synctest.Wait() + select { + case result := <-resultChan: + require.NoError(t, result.err) + require.Equal(t, []byte("foobar"), result.data) + default: + t.Fatal("should have received a datagram frame") + } + + // unblock when the context is canceled + ctx, cancel := context.WithCancel(context.Background()) + errChan := make(chan error, 1) + go func() { + _, err := queue.Receive(ctx) + errChan <- err + }() + + synctest.Wait() + select { + case <-errChan: + t.Fatal("expected to not receive error") + default: + } + + cancel() + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, context.Canceled) + default: + t.Fatal("should have received a context canceled error") + } + }) +} + +func TestDatagramQueueClose(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + queue := newDatagramQueue(func() {}, utils.DefaultLogger) + + for range maxDatagramSendQueueLen { + require.NoError(t, queue.Add(&wire.DatagramFrame{Data: []byte{0}})) + } + errChan1 := make(chan error, 1) + go func() { errChan1 <- queue.Add(&wire.DatagramFrame{Data: []byte("foobar")}) }() + errChan2 := make(chan error, 1) + go func() { + _, err := queue.Receive(context.Background()) + errChan2 <- err + }() + + queue.CloseWithError(assert.AnError) + synctest.Wait() + + select { + case err := <-errChan1: + require.ErrorIs(t, err, assert.AnError) + default: + t.Fatal("should have received an error") + } + + select { + case err := <-errChan2: + require.ErrorIs(t, err, assert.AnError) + default: + t.Fatal("should have received an error") + } + }) +} diff --git a/third_party/quic-go/errors.go b/third_party/quic-go/errors.go new file mode 100644 index 0000000..c1c9c3e --- /dev/null +++ b/third_party/quic-go/errors.go @@ -0,0 +1,107 @@ +package quic + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/qerr" +) + +type ( + // TransportError indicates an error that occurred on the QUIC transport layer. + // Every transport error other than CONNECTION_REFUSED and APPLICATION_ERROR is + // likely a bug in the implementation. + TransportError = qerr.TransportError + // ApplicationError is an application-defined error. + ApplicationError = qerr.ApplicationError + // VersionNegotiationError indicates a failure to negotiate a QUIC version. + VersionNegotiationError = qerr.VersionNegotiationError + // StatelessResetError indicates a stateless reset was received. + // This can happen when the peer reboots, or when packets are misrouted. + // See section 10.3 of RFC 9000 for details. + StatelessResetError = qerr.StatelessResetError + // IdleTimeoutError indicates that the connection timed out because it was inactive for too long. + IdleTimeoutError = qerr.IdleTimeoutError + // HandshakeTimeoutError indicates that the connection timed out before completing the handshake. + HandshakeTimeoutError = qerr.HandshakeTimeoutError +) + +type ( + // TransportErrorCode is a QUIC transport error code, see section 20 of RFC 9000. + TransportErrorCode = qerr.TransportErrorCode + // ApplicationErrorCode is an QUIC application error code. + ApplicationErrorCode = qerr.ApplicationErrorCode + // StreamErrorCode is a QUIC stream error code. The meaning of the value is defined by the application. + StreamErrorCode = qerr.StreamErrorCode +) + +const ( + // NoError is the NO_ERROR transport error code. + NoError = qerr.NoError + // InternalError is the INTERNAL_ERROR transport error code. + InternalError = qerr.InternalError + // ConnectionRefused is the CONNECTION_REFUSED transport error code. + ConnectionRefused = qerr.ConnectionRefused + // FlowControlError is the FLOW_CONTROL_ERROR transport error code. + FlowControlError = qerr.FlowControlError + // StreamLimitError is the STREAM_LIMIT_ERROR transport error code. + StreamLimitError = qerr.StreamLimitError + // StreamStateError is the STREAM_STATE_ERROR transport error code. + StreamStateError = qerr.StreamStateError + // FinalSizeError is the FINAL_SIZE_ERROR transport error code. + FinalSizeError = qerr.FinalSizeError + // FrameEncodingError is the FRAME_ENCODING_ERROR transport error code. + FrameEncodingError = qerr.FrameEncodingError + // TransportParameterError is the TRANSPORT_PARAMETER_ERROR transport error code. + TransportParameterError = qerr.TransportParameterError + // ConnectionIDLimitError is the CONNECTION_ID_LIMIT_ERROR transport error code. + ConnectionIDLimitError = qerr.ConnectionIDLimitError + // ProtocolViolation is the PROTOCOL_VIOLATION transport error code. + ProtocolViolation = qerr.ProtocolViolation + // InvalidToken is the INVALID_TOKEN transport error code. + InvalidToken = qerr.InvalidToken + // ApplicationErrorErrorCode is the APPLICATION_ERROR transport error code. + ApplicationErrorErrorCode = qerr.ApplicationErrorErrorCode + // CryptoBufferExceeded is the CRYPTO_BUFFER_EXCEEDED transport error code. + CryptoBufferExceeded = qerr.CryptoBufferExceeded + // KeyUpdateError is the KEY_UPDATE_ERROR transport error code. + KeyUpdateError = qerr.KeyUpdateError + // AEADLimitReached is the AEAD_LIMIT_REACHED transport error code. + AEADLimitReached = qerr.AEADLimitReached + // NoViablePathError is the NO_VIABLE_PATH_ERROR transport error code. + NoViablePathError = qerr.NoViablePathError +) + +// A StreamError is used to signal stream cancellations. +// It is returned from the Read and Write methods of the [ReceiveStream], [SendStream] and [Stream]. +type StreamError struct { + StreamID StreamID + ErrorCode StreamErrorCode + Remote bool +} + +func (e *StreamError) Is(target error) bool { + t, ok := target.(*StreamError) + return ok && e.StreamID == t.StreamID && e.ErrorCode == t.ErrorCode && e.Remote == t.Remote +} + +func (e *StreamError) Error() string { + pers := "local" + if e.Remote { + pers = "remote" + } + return fmt.Sprintf("stream %d canceled by %s with error code %d", e.StreamID, pers, e.ErrorCode) +} + +// DatagramTooLargeError is returned from Conn.SendDatagram if the payload is too large to be sent. +type DatagramTooLargeError struct { + MaxDatagramPayloadSize int64 +} + +func (e *DatagramTooLargeError) Is(target error) bool { + t, ok := target.(*DatagramTooLargeError) + return ok && e.MaxDatagramPayloadSize == t.MaxDatagramPayloadSize +} + +func (e *DatagramTooLargeError) Error() string { + return fmt.Sprintf("DATAGRAM frame too large (maximum: %d bytes)", e.MaxDatagramPayloadSize) +} diff --git a/third_party/quic-go/errors_test.go b/third_party/quic-go/errors_test.go new file mode 100644 index 0000000..870b37e --- /dev/null +++ b/third_party/quic-go/errors_test.go @@ -0,0 +1,37 @@ +package quic + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestStreamError(t *testing.T) { + require.True(t, errors.Is( + &StreamError{StreamID: 1, ErrorCode: 2, Remote: true}, + &StreamError{StreamID: 1, ErrorCode: 2, Remote: true}, + )) + require.False(t, errors.Is(&StreamError{StreamID: 1}, &StreamError{StreamID: 2})) + require.False(t, errors.Is(&StreamError{StreamID: 1}, &StreamError{StreamID: 2})) + require.Equal(t, + "stream 1 canceled by remote with error code 2", + (&StreamError{StreamID: 1, ErrorCode: 2, Remote: true}).Error(), + ) + require.Equal(t, + "stream 42 canceled by local with error code 1337", + (&StreamError{StreamID: 42, ErrorCode: 1337, Remote: false}).Error(), + ) +} + +func TestDatagramTooLargeError(t *testing.T) { + require.True(t, errors.Is( + &DatagramTooLargeError{MaxDatagramPayloadSize: 1024}, + &DatagramTooLargeError{MaxDatagramPayloadSize: 1024}, + )) + require.False(t, errors.Is( + &DatagramTooLargeError{MaxDatagramPayloadSize: 1024}, + &DatagramTooLargeError{MaxDatagramPayloadSize: 1025}, + )) + require.Equal(t, "DATAGRAM frame too large (maximum: 1024 bytes)", (&DatagramTooLargeError{MaxDatagramPayloadSize: 1024}).Error()) +} diff --git a/third_party/quic-go/example/client/main.go b/third_party/quic-go/example/client/main.go new file mode 100644 index 0000000..0bd6917 --- /dev/null +++ b/third_party/quic-go/example/client/main.go @@ -0,0 +1,81 @@ +package main + +import ( + "bytes" + "crypto/tls" + "crypto/x509" + "flag" + "io" + "log" + "net/http" + "os" + "sync" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/internal/testdata" +) + +func main() { + quiet := flag.Bool("q", false, "don't print the data") + keyLogFile := flag.String("keylog", "", "key log file") + insecure := flag.Bool("insecure", false, "skip certificate verification") + flag.Parse() + urls := flag.Args() + + var keyLog io.Writer + if len(*keyLogFile) > 0 { + f, err := os.Create(*keyLogFile) + if err != nil { + log.Fatal(err) + } + defer f.Close() + keyLog = f + } + + pool, err := x509.SystemCertPool() + if err != nil { + log.Fatal(err) + } + testdata.AddRootCA(pool) + + roundTripper := &http3.Transport{ + TLSClientConfig: &tls.Config{ + RootCAs: pool, + InsecureSkipVerify: *insecure, + KeyLogWriter: keyLog, + }, + QUICConfig: &quic.Config{ + Tracer: qlog.DefaultConnectionTracer, + }, + } + defer roundTripper.Close() + hclient := &http.Client{ + Transport: roundTripper, + } + + var wg sync.WaitGroup + for _, addr := range urls { + log.Printf("GET %s", addr) + wg.Go(func() { + rsp, err := hclient.Get(addr) + if err != nil { + log.Fatal(err) + } + log.Printf("Got response for %s: %#v", addr, rsp) + + body := &bytes.Buffer{} + _, err = io.Copy(body, rsp.Body) + if err != nil { + log.Fatal(err) + } + if *quiet { + log.Printf("Response Body: %d bytes", body.Len()) + } else { + log.Printf("Response Body (%d bytes):\n%s", body.Len(), body.Bytes()) + } + }) + } + wg.Wait() +} diff --git a/third_party/quic-go/example/echo/echo.go b/third_party/quic-go/example/echo/echo.go new file mode 100644 index 0000000..cf02fd2 --- /dev/null +++ b/third_party/quic-go/example/echo/echo.go @@ -0,0 +1,113 @@ +package main + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "fmt" + "io" + "log" + "math/big" + + "github.com/apernet/quic-go" +) + +const addr = "localhost:4242" + +const message = "foobar" + +// We start a server echoing data on the first stream the client opens, +// then connect with a client, send the message, and wait for its receipt. +func main() { + go func() { log.Fatal(echoServer()) }() + + if err := clientMain(); err != nil { + panic(err) + } +} + +// Start a server that echos all data on the first stream opened by the client +func echoServer() error { + listener, err := quic.ListenAddr(addr, generateTLSConfig(), nil) + if err != nil { + return err + } + defer listener.Close() + + conn, err := listener.Accept(context.Background()) + if err != nil { + return err + } + + stream, err := conn.AcceptStream(context.Background()) + if err != nil { + panic(err) + } + defer stream.Close() + + // Echo through the loggingWriter + _, err = io.Copy(loggingWriter{stream}, stream) + return err +} + +func clientMain() error { + tlsConf := &tls.Config{ + InsecureSkipVerify: true, + NextProtos: []string{"quic-echo-example"}, + } + conn, err := quic.DialAddr(context.Background(), addr, tlsConf, nil) + if err != nil { + return err + } + defer conn.CloseWithError(0, "") + + stream, err := conn.OpenStreamSync(context.Background()) + if err != nil { + return err + } + defer stream.Close() + + fmt.Printf("Client: Sending '%s'\n", message) + if _, err := stream.Write([]byte(message)); err != nil { + return err + } + + buf := make([]byte, len(message)) + if _, err := io.ReadFull(stream, buf); err != nil { + return err + } + fmt.Printf("Client: Got '%s'\n", buf) + + return nil +} + +// A wrapper for io.Writer that also logs the message. +type loggingWriter struct{ io.Writer } + +func (w loggingWriter) Write(b []byte) (int, error) { + fmt.Printf("Server: Got '%s'\n", string(b)) + return w.Writer.Write(b) +} + +// Setup a bare-bones TLS config for the server +func generateTLSConfig() *tls.Config { + _, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + panic(err) + } + template := x509.Certificate{SerialNumber: big.NewInt(1)} + certDER, err := x509.CreateCertificate(rand.Reader, &template, &template, priv.Public(), priv) + if err != nil { + panic(err) + } + + return &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: [][]byte{certDER}, + PrivateKey: priv, + }}, + NextProtos: []string{"quic-echo-example"}, + } +} diff --git a/third_party/quic-go/example/main.go b/third_party/quic-go/example/main.go new file mode 100644 index 0000000..a95603a --- /dev/null +++ b/third_party/quic-go/example/main.go @@ -0,0 +1,182 @@ +package main + +import ( + "crypto/md5" + "errors" + "flag" + "fmt" + "io" + "log" + "mime/multipart" + "net/http" + "strconv" + "strings" + "sync" + + _ "net/http/pprof" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/internal/testdata" +) + +type binds []string + +func (b binds) String() string { + return strings.Join(b, ",") +} + +func (b *binds) Set(v string) error { + *b = strings.Split(v, ",") + return nil +} + +// Size is needed by the /demo/upload handler to determine the size of the uploaded file +type Size interface { + Size() int64 +} + +// See https://en.wikipedia.org/wiki/Lehmer_random_number_generator +func generatePRData(l int) []byte { + res := make([]byte, l) + seed := uint64(1) + for i := range l { + seed = seed * 48271 % 2147483647 + res[i] = byte(seed) + } + return res +} + +func setupHandler(www string) http.Handler { + mux := http.NewServeMux() + + if len(www) > 0 { + mux.Handle("/", http.FileServer(http.Dir(www))) + } else { + mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + fmt.Printf("%#v\n", r) + const maxSize = 1 << 30 // 1 GB + num, err := strconv.ParseInt(strings.ReplaceAll(r.RequestURI, "/", ""), 10, 64) + if err != nil || num <= 0 || num > maxSize { + w.WriteHeader(400) + return + } + w.Write(generatePRData(int(num))) + }) + } + + mux.HandleFunc("/demo/tile", func(w http.ResponseWriter, r *http.Request) { + // Small 40x40 png + w.Write([]byte{ + 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, + 0x49, 0x48, 0x44, 0x52, 0x00, 0x00, 0x00, 0x28, 0x00, 0x00, 0x00, 0x28, + 0x01, 0x03, 0x00, 0x00, 0x00, 0xb6, 0x30, 0x2a, 0x2e, 0x00, 0x00, 0x00, + 0x03, 0x50, 0x4c, 0x54, 0x45, 0x5a, 0xc3, 0x5a, 0xad, 0x38, 0xaa, 0xdb, + 0x00, 0x00, 0x00, 0x0b, 0x49, 0x44, 0x41, 0x54, 0x78, 0x01, 0x63, 0x18, + 0x61, 0x00, 0x00, 0x00, 0xf0, 0x00, 0x01, 0xe2, 0xb8, 0x75, 0x22, 0x00, + 0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae, 0x42, 0x60, 0x82, + }) + }) + + mux.HandleFunc("/demo/tiles", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "") + for i := range 200 { + fmt.Fprintf(w, ``, i) + } + io.WriteString(w, "") + }) + + mux.HandleFunc("/demo/echo", func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + if err != nil { + fmt.Printf("error reading body while handling /echo: %s\n", err.Error()) + } + w.Write(body) + }) + + // accept file uploads and return the MD5 of the uploaded file + // maximum accepted file size is 1 GB + mux.HandleFunc("/demo/upload", func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost { + err := r.ParseMultipartForm(1 << 30) // 1 GB + if err == nil { + var file multipart.File + file, _, err = r.FormFile("uploadfile") + if err == nil { + var size int64 + if sizeInterface, ok := file.(Size); ok { + size = sizeInterface.Size() + b := make([]byte, size) + file.Read(b) + md5 := md5.Sum(b) + fmt.Fprintf(w, "%x", md5) + return + } + err = errors.New("couldn't get uploaded file size") + } + } + log.Printf("Error receiving upload: %#v", err) + } + io.WriteString(w, `

+
+ + `) + }) + + return mux +} + +func main() { + // defer profile.Start().Stop() + go func() { + log.Println(http.ListenAndServe("localhost:6060", nil)) + }() + // runtime.SetBlockProfileRate(1) + + bs := binds{} + flag.Var(&bs, "bind", "bind to") + www := flag.String("www", "", "www data") + tcp := flag.Bool("tcp", false, "also listen on TCP") + key := flag.String("key", "", "TLS key (requires -cert option)") + cert := flag.String("cert", "", "TLS certificate (requires -key option)") + flag.Parse() + + if len(bs) == 0 { + bs = binds{"localhost:6121"} + } + + handler := setupHandler(*www) + + var wg sync.WaitGroup + var certFile, keyFile string + if *key != "" && *cert != "" { + keyFile = *key + certFile = *cert + } else { + certFile, keyFile = testdata.GetCertificatePaths() + } + for _, b := range bs { + fmt.Println("listening on", b) + bCap := b + wg.Go(func() { + var err error + if *tcp { + err = http3.ListenAndServeTLS(bCap, certFile, keyFile, handler) + } else { + server := http3.Server{ + Handler: handler, + Addr: bCap, + QUICConfig: &quic.Config{ + Tracer: qlog.DefaultConnectionTracer, + }, + } + err = server.ListenAndServeTLS(certFile, keyFile) + } + if err != nil { + fmt.Println(err) + } + }) + } + wg.Wait() +} diff --git a/third_party/quic-go/flow_controller_base.go b/third_party/quic-go/flow_controller_base.go new file mode 100644 index 0000000..31ba98a --- /dev/null +++ b/third_party/quic-go/flow_controller_base.go @@ -0,0 +1,84 @@ +package quic + +import ( + "sync" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +type receiveFlowController struct { + //nolint:structcheck // The mutex is used both by the stream and the connection flow controller + mutex sync.Mutex + bytesRead protocol.ByteCount + highestReceived protocol.ByteCount + receiveWindow protocol.ByteCount + receiveWindowSize protocol.ByteCount + maxReceiveWindowSize protocol.ByteCount + + allowWindowIncrease func(size protocol.ByteCount) bool + + epochStartTime monotime.Time + epochStartOffset protocol.ByteCount + rttStats *utils.RTTStats + + logger utils.Logger +} + +// needs to be called with locked mutex +func (c *receiveFlowController) addBytesRead(n protocol.ByteCount) { + c.bytesRead += n +} + +func (c *receiveFlowController) hasWindowUpdate() bool { + bytesRemaining := c.receiveWindow - c.bytesRead + // update the window when more than the threshold was consumed + return bytesRemaining <= protocol.ByteCount(float64(c.receiveWindowSize)*(1-protocol.WindowUpdateThreshold)) +} + +// getWindowUpdate updates the receive window, if necessary +// it returns the new offset +func (c *receiveFlowController) getWindowUpdate(now monotime.Time) protocol.ByteCount { + if !c.hasWindowUpdate() { + return 0 + } + + c.maybeAdjustWindowSize(now) + c.receiveWindow = c.bytesRead + c.receiveWindowSize + return c.receiveWindow +} + +// maybeAdjustWindowSize increases the receiveWindowSize if we're sending updates too often. +// For details about auto-tuning, see https://docs.google.com/document/d/1SExkMmGiz8VYzV3s9E35JQlJ73vhzCekKkDi85F1qCE/edit?usp=sharing. +func (c *receiveFlowController) maybeAdjustWindowSize(now monotime.Time) { + bytesReadInEpoch := c.bytesRead - c.epochStartOffset + // don't do anything if less than half the window has been consumed + if bytesReadInEpoch <= c.receiveWindowSize/2 { + return + } + rtt := c.rttStats.SmoothedRTT() + if rtt == 0 { + return + } + + fraction := float64(bytesReadInEpoch) / float64(c.receiveWindowSize) + if now.Sub(c.epochStartTime) < time.Duration(4*fraction*float64(rtt)) { + // window is consumed too fast, try to increase the window size + newSize := min(2*c.receiveWindowSize, c.maxReceiveWindowSize) + if newSize > c.receiveWindowSize && (c.allowWindowIncrease == nil || c.allowWindowIncrease(newSize-c.receiveWindowSize)) { + c.receiveWindowSize = newSize + } + } + c.startNewAutoTuningEpoch(now) +} + +func (c *receiveFlowController) startNewAutoTuningEpoch(now monotime.Time) { + c.epochStartTime = now + c.epochStartOffset = c.bytesRead +} + +func (c *receiveFlowController) checkFlowControlViolation() bool { + return c.highestReceived > c.receiveWindow +} diff --git a/third_party/quic-go/flow_controller_connection.go b/third_party/quic-go/flow_controller_connection.go new file mode 100644 index 0000000..beddd67 --- /dev/null +++ b/third_party/quic-go/flow_controller_connection.go @@ -0,0 +1,186 @@ +package quic + +import ( + "errors" + "fmt" + "sync" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" +) + +type connectionFlowController struct { + receiveFlowController + + // Protects send-side state, which TryWriteAll can access from application goroutines. + sendMutex sync.Mutex + bytesSent protocol.ByteCount + sendWindow protocol.ByteCount + lastBlockedAt protocol.ByteCount +} + +// newConnectionFlowController gets a new flow controller for the connection. +// It is created before we receive the peer's transport parameters, thus it starts with a sendWindow of 0. +func newConnectionFlowController( + receiveWindow protocol.ByteCount, + maxReceiveWindow protocol.ByteCount, + allowWindowIncrease func(size protocol.ByteCount) bool, + rttStats *utils.RTTStats, + logger utils.Logger, +) *connectionFlowController { + return &connectionFlowController{ + receiveFlowController: receiveFlowController{ + rttStats: rttStats, + receiveWindow: receiveWindow, + receiveWindowSize: receiveWindow, + maxReceiveWindowSize: maxReceiveWindow, + allowWindowIncrease: allowWindowIncrease, + logger: logger, + }, + } +} + +// IncrementHighestReceived adds an increment to the highestReceived value +func (c *connectionFlowController) IncrementHighestReceived(increment protocol.ByteCount, now monotime.Time) error { + c.mutex.Lock() + defer c.mutex.Unlock() + + // If this is the first frame received on this connection, start flow-control auto-tuning. + if c.highestReceived == 0 { + c.startNewAutoTuningEpoch(now) + } + c.highestReceived += increment + + if c.checkFlowControlViolation() { + return &qerr.TransportError{ + ErrorCode: qerr.FlowControlError, + ErrorMessage: fmt.Sprintf("received %d bytes for the connection, allowed %d bytes", c.highestReceived, c.receiveWindow), + } + } + return nil +} + +func (c *connectionFlowController) AddBytesRead(n protocol.ByteCount) (hasWindowUpdate bool) { + c.mutex.Lock() + defer c.mutex.Unlock() + + c.addBytesRead(n) + return c.hasWindowUpdate() +} + +// TryAddBytesSent adds n bytes if sufficient connection-level send credit is available. +func (c *connectionFlowController) TryAddBytesSent(n protocol.ByteCount) bool { + c.sendMutex.Lock() + defer c.sendMutex.Unlock() + + if c.bytesSent > c.sendWindow || n > c.sendWindow-c.bytesSent { + return false + } + c.bytesSent += n + return true +} + +// AddBytesSentWithLimiter adds the limiter-approved portion of the available connection-level send credit. +func (c *connectionFlowController) AddBytesSentWithLimiter( + n protocol.ByteCount, + limiter func(int) int, +) (protocol.ByteCount, bool) { + c.sendMutex.Lock() + defer c.sendMutex.Unlock() + + if c.bytesSent >= c.sendWindow { + return 0, false + } + n = min(n, c.sendWindow-c.bytesSent) + added := min( + max(protocol.ByteCount(limiter(int(n))), 0), + n, + ) + c.bytesSent += added + return added, added < n +} + +// UpdateSendWindow is called after receiving a MAX_DATA frame. +func (c *connectionFlowController) UpdateSendWindow(offset protocol.ByteCount) (updated bool) { + c.sendMutex.Lock() + defer c.sendMutex.Unlock() + + if offset > c.sendWindow { + c.sendWindow = offset + return true + } + return false +} + +func (c *connectionFlowController) SendWindowSize() protocol.ByteCount { + c.sendMutex.Lock() + defer c.sendMutex.Unlock() + + return c.sendWindow - c.bytesSent +} + +// IsNewlyBlocked says if it is newly blocked by connection flow control. +// For every offset, it only returns true once. +// If it is blocked, the offset is returned. +func (c *connectionFlowController) IsNewlyBlocked() (bool, protocol.ByteCount) { + c.sendMutex.Lock() + defer c.sendMutex.Unlock() + + if c.bytesSent < c.sendWindow || c.sendWindow == c.lastBlockedAt { + return false, 0 + } + c.lastBlockedAt = c.sendWindow + return true, c.sendWindow +} + +func (c *connectionFlowController) GetWindowUpdate(now monotime.Time) protocol.ByteCount { + c.mutex.Lock() + defer c.mutex.Unlock() + + oldWindowSize := c.receiveWindowSize + offset := c.getWindowUpdate(now) + if c.logger.Debug() && oldWindowSize < c.receiveWindowSize { + c.logger.Debugf("Increasing receive flow control window for the connection to %d kB", c.receiveWindowSize/(1<<10)) + } + return offset +} + +// EnsureMinimumWindowSize sets a minimum window size +// it should make sure that the connection-level window is increased when a stream-level window grows +func (c *connectionFlowController) EnsureMinimumWindowSize(inc protocol.ByteCount, now monotime.Time) { + c.mutex.Lock() + defer c.mutex.Unlock() + + if inc <= c.receiveWindowSize { + return + } + newSize := min(inc, c.maxReceiveWindowSize) + if delta := newSize - c.receiveWindowSize; delta > 0 && c.allowWindowIncrease(delta) { + c.receiveWindowSize = newSize + if c.logger.Debug() { + c.logger.Debugf("Increasing receive flow control window for the connection to %d, in response to stream flow control window increase", newSize) + } + } + c.startNewAutoTuningEpoch(now) +} + +// Reset rests the flow controller. This happens when 0-RTT is rejected. +// All stream data is invalidated, it's as if we had never opened a stream and never sent any data. +// At that point, we only have sent stream data, but we didn't have the keys to open 1-RTT keys yet. +func (c *connectionFlowController) Reset() error { + c.mutex.Lock() + defer c.mutex.Unlock() + + if c.bytesRead > 0 || c.highestReceived > 0 || !c.epochStartTime.IsZero() { + return errors.New("flow controller reset after reading data") + } + c.sendMutex.Lock() + defer c.sendMutex.Unlock() + + c.bytesSent = 0 + c.lastBlockedAt = 0 + c.sendWindow = 0 + return nil +} diff --git a/third_party/quic-go/flow_controller_connection_test.go b/third_party/quic-go/flow_controller_connection_test.go new file mode 100644 index 0000000..43ebcd4 --- /dev/null +++ b/third_party/quic-go/flow_controller_connection_test.go @@ -0,0 +1,77 @@ +package quic + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + + "github.com/stretchr/testify/require" +) + +func TestConnectionFlowControlWindowUpdate(t *testing.T) { + fc := newConnectionFlowController( + 100, // initial receive window + 100, // max receive window + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + require.False(t, fc.AddBytesRead(1)) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) + require.True(t, fc.AddBytesRead(99)) + require.Equal(t, protocol.ByteCount(200), fc.GetWindowUpdate(monotime.Now())) +} + +func TestConnectionWindowAutoTuningNotAllowed(t *testing.T) { + // the RTT is 1 second + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(time.Second, 0) + require.Equal(t, time.Second, rttStats.SmoothedRTT()) + + callbackCalledWith := protocol.InvalidByteCount + fc := newConnectionFlowController( + 100, // initial receive window + 150, // max receive window + func(size protocol.ByteCount) bool { + callbackCalledWith = size + return false + }, + rttStats, + utils.DefaultLogger, + ) + now := monotime.Now() + require.NoError(t, fc.IncrementHighestReceived(100, now)) + fc.AddBytesRead(90) + require.Equal(t, protocol.InvalidByteCount, callbackCalledWith) + require.Equal(t, protocol.ByteCount(90+100), fc.GetWindowUpdate(now.Add(time.Millisecond))) + require.Equal(t, protocol.ByteCount(150-100), callbackCalledWith) +} + +func TestConnectionFlowControlViolation(t *testing.T) { + fc := newConnectionFlowController(100, 100, nil, utils.NewRTTStats(), utils.DefaultLogger) + require.NoError(t, fc.IncrementHighestReceived(40, monotime.Now())) + require.NoError(t, fc.IncrementHighestReceived(60, monotime.Now())) + err := fc.IncrementHighestReceived(1, monotime.Now()) + var terr *qerr.TransportError + require.ErrorAs(t, err, &terr) + require.Equal(t, qerr.FlowControlError, terr.ErrorCode) +} + +func TestConnectionFlowControllerReset(t *testing.T) { + fc := newConnectionFlowController(0, 0, nil, utils.NewRTTStats(), utils.DefaultLogger) + fc.UpdateSendWindow(100) + require.True(t, fc.TryAddBytesSent(10)) + require.Equal(t, protocol.ByteCount(90), fc.SendWindowSize()) + require.NoError(t, fc.Reset()) + require.Zero(t, fc.SendWindowSize()) +} + +func TestConnectionFlowControllerResetAfterReading(t *testing.T) { + fc := newConnectionFlowController(0, 0, nil, utils.NewRTTStats(), utils.DefaultLogger) + fc.AddBytesRead(1) + require.EqualError(t, fc.Reset(), "flow controller reset after reading data") +} diff --git a/third_party/quic-go/flow_controller_stream.go b/third_party/quic-go/flow_controller_stream.go new file mode 100644 index 0000000..2cb7531 --- /dev/null +++ b/third_party/quic-go/flow_controller_stream.go @@ -0,0 +1,196 @@ +package quic + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" +) + +type streamFlowController struct { + receiveFlowController + + bytesSent protocol.ByteCount + sendWindow protocol.ByteCount + lastBlockedAt protocol.ByteCount + + streamID protocol.StreamID + + connection *connectionFlowController + + receivedFinalOffset bool +} + +// newStreamFlowController gets a new flow controller for a stream. +func newStreamFlowController( + streamID protocol.StreamID, + cfc *connectionFlowController, + receiveWindow protocol.ByteCount, + maxReceiveWindow protocol.ByteCount, + initialSendWindow protocol.ByteCount, + rttStats *utils.RTTStats, + logger utils.Logger, +) *streamFlowController { + return &streamFlowController{ + streamID: streamID, + connection: cfc, + sendWindow: initialSendWindow, + receiveFlowController: receiveFlowController{ + rttStats: rttStats, + receiveWindow: receiveWindow, + receiveWindowSize: receiveWindow, + maxReceiveWindowSize: maxReceiveWindow, + logger: logger, + }, + } +} + +// UpdateHighestReceived updates the highestReceived value, if the offset is higher. +func (c *streamFlowController) UpdateHighestReceived(offset protocol.ByteCount, final bool, now monotime.Time) error { + // If the final offset for this stream is already known, check for consistency. + if c.receivedFinalOffset { + // If we receive another final offset, check that it's the same. + if final && offset != c.highestReceived { + return &qerr.TransportError{ + ErrorCode: qerr.FinalSizeError, + ErrorMessage: fmt.Sprintf("received inconsistent final offset for stream %d (old: %d, new: %d bytes)", c.streamID, c.highestReceived, offset), + } + } + // Check that the offset is below the final offset. + if offset > c.highestReceived { + return &qerr.TransportError{ + ErrorCode: qerr.FinalSizeError, + ErrorMessage: fmt.Sprintf("received offset %d for stream %d, but final offset was already received at %d", offset, c.streamID, c.highestReceived), + } + } + } + + if final { + c.receivedFinalOffset = true + } + if offset == c.highestReceived { + return nil + } + // A higher offset was received before. This can happen due to reordering. + if offset < c.highestReceived { + if final { + return &qerr.TransportError{ + ErrorCode: qerr.FinalSizeError, + ErrorMessage: fmt.Sprintf("received final offset %d for stream %d, but already received offset %d before", offset, c.streamID, c.highestReceived), + } + } + return nil + } + + // If this is the first frame received for this stream, start flow-control auto-tuning. + if c.highestReceived == 0 { + c.startNewAutoTuningEpoch(now) + } + increment := offset - c.highestReceived + c.highestReceived = offset + + if c.checkFlowControlViolation() { + return &qerr.TransportError{ + ErrorCode: qerr.FlowControlError, + ErrorMessage: fmt.Sprintf("received %d bytes on stream %d, allowed %d bytes", offset, c.streamID, c.receiveWindow), + } + } + return c.connection.IncrementHighestReceived(increment, now) +} + +func (c *streamFlowController) AddBytesRead(n protocol.ByteCount) (hasStreamWindowUpdate, hasConnWindowUpdate bool) { + c.mutex.Lock() + c.addBytesRead(n) + hasStreamWindowUpdate = c.shouldQueueWindowUpdate() + c.mutex.Unlock() + hasConnWindowUpdate = c.connection.AddBytesRead(n) + return +} + +func (c *streamFlowController) Abandon() { + c.mutex.Lock() + unread := c.highestReceived - c.bytesRead + c.bytesRead = c.highestReceived + c.mutex.Unlock() + if unread > 0 { + c.connection.AddBytesRead(unread) + } +} + +func (c *streamFlowController) UpdateSendWindow(offset protocol.ByteCount) (updated bool) { + if offset > c.sendWindow { + c.sendWindow = offset + return true + } + return false +} + +// TryAddBytesSent adds n bytes if sufficient stream- and connection-level send credit is available. +func (c *streamFlowController) TryAddBytesSent(n protocol.ByteCount) bool { + if c.bytesSent > c.sendWindow || n > c.sendWindow-c.bytesSent { + return false + } + if !c.connection.TryAddBytesSent(n) { + return false + } + c.bytesSent += n + return true +} + +// AddBytesSentWithLimiter adds the limiter-approved portion of the available stream- and connection-level send credit. +func (c *streamFlowController) AddBytesSentWithLimiter( + n protocol.ByteCount, + limiter func(int) int, +) (protocol.ByteCount, bool) { + if c.bytesSent >= c.sendWindow { + return 0, false + } + n = min(n, c.sendWindow-c.bytesSent) + added, limited := c.connection.AddBytesSentWithLimiter(n, limiter) + c.bytesSent += added + return added, limited +} + +func (c *streamFlowController) SendWindowSize() protocol.ByteCount { + return min(c.sendWindow-c.bytesSent, c.connection.SendWindowSize()) +} + +func (c *streamFlowController) IsNewlyBlocked() bool { + blocked, _ := c.isNewlyBlocked() + return blocked +} + +func (c *streamFlowController) isNewlyBlocked() (bool, protocol.ByteCount) { + if c.bytesSent < c.sendWindow || c.sendWindow == c.lastBlockedAt { + return false, 0 + } + c.lastBlockedAt = c.sendWindow + return true, c.sendWindow +} + +func (c *streamFlowController) shouldQueueWindowUpdate() bool { + return !c.receivedFinalOffset && c.hasWindowUpdate() +} + +func (c *streamFlowController) GetWindowUpdate(now monotime.Time) protocol.ByteCount { + // If we already received the final offset for this stream, the peer won't need any additional flow control credit. + if c.receivedFinalOffset { + return 0 + } + + c.mutex.Lock() + defer c.mutex.Unlock() + + oldWindowSize := c.receiveWindowSize + offset := c.getWindowUpdate(now) + if c.receiveWindowSize > oldWindowSize { // auto-tuning enlarged the window size + c.logger.Debugf("Increasing receive flow control window for stream %d to %d", c.streamID, c.receiveWindowSize) + c.connection.EnsureMinimumWindowSize( + protocol.ByteCount(float64(c.receiveWindowSize)*protocol.ConnectionFlowControlMultiplier), + now, + ) + } + return offset +} diff --git a/third_party/quic-go/flow_controller_stream_test.go b/third_party/quic-go/flow_controller_stream_test.go new file mode 100644 index 0000000..f8d52c0 --- /dev/null +++ b/third_party/quic-go/flow_controller_stream_test.go @@ -0,0 +1,319 @@ +package quic + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + + "github.com/stretchr/testify/require" +) + +func TestStreamFlowControlReceiving(t *testing.T) { + fc := newStreamFlowController( + 42, + newConnectionFlowController( + protocol.MaxByteCount, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ), + 100, + protocol.MaxByteCount, + protocol.MaxByteCount, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + + require.NoError(t, fc.UpdateHighestReceived(50, false, monotime.Now())) + // duplicates are fine + require.NoError(t, fc.UpdateHighestReceived(50, false, monotime.Now())) + // reordering is fine + require.NoError(t, fc.UpdateHighestReceived(40, false, monotime.Now())) + require.NoError(t, fc.UpdateHighestReceived(60, false, monotime.Now())) + + // exceeding the limit is not fine + err := fc.UpdateHighestReceived(101, false, monotime.Now()) + var terr *qerr.TransportError + require.ErrorAs(t, err, &terr) + require.Equal(t, qerr.FlowControlError, terr.ErrorCode) + require.Equal(t, "received 101 bytes on stream 42, allowed 100 bytes", terr.ErrorMessage) +} + +func TestStreamFlowControllerFinalOffset(t *testing.T) { + newFC := func() *streamFlowController { + return newStreamFlowController( + 42, + newConnectionFlowController( + protocol.MaxByteCount, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ), + protocol.MaxByteCount, + protocol.MaxByteCount, + protocol.MaxByteCount, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + } + + t.Run("duplicate final offset", func(t *testing.T) { + fc := newFC() + require.NoError(t, fc.UpdateHighestReceived(50, true, monotime.Now())) + // it is valid to receive the same final offset multiple times + require.NoError(t, fc.UpdateHighestReceived(50, true, monotime.Now())) + }) + + t.Run("inconsistent final offset", func(t *testing.T) { + fc := newFC() + require.NoError(t, fc.UpdateHighestReceived(50, true, monotime.Now())) + err := fc.UpdateHighestReceived(51, true, monotime.Now()) + require.Error(t, err) + var terr *qerr.TransportError + require.ErrorAs(t, err, &terr) + require.Equal(t, qerr.FinalSizeError, terr.ErrorCode) + require.Equal(t, "received inconsistent final offset for stream 42 (old: 50, new: 51 bytes)", terr.ErrorMessage) + }) + + t.Run("non-final offset past final offset", func(t *testing.T) { + fc := newFC() + require.NoError(t, fc.UpdateHighestReceived(50, true, monotime.Now())) + // No matter the ordering, it's never ok to receive an offset past the final offset. + err := fc.UpdateHighestReceived(60, false, monotime.Now()) + var terr *qerr.TransportError + require.ErrorAs(t, err, &terr) + require.Equal(t, qerr.FinalSizeError, terr.ErrorCode) + require.Equal(t, "received offset 60 for stream 42, but final offset was already received at 50", terr.ErrorMessage) + }) + + t.Run("final offset smaller than previous offset", func(t *testing.T) { + fc := newFC() + require.NoError(t, fc.UpdateHighestReceived(50, false, monotime.Now())) + // If we received offset already, it's invalid to receive a smaller final offset. + err := fc.UpdateHighestReceived(40, true, monotime.Now()) + var terr *qerr.TransportError + require.ErrorAs(t, err, &terr) + require.Equal(t, qerr.FinalSizeError, terr.ErrorCode) + require.Equal(t, "received final offset 40 for stream 42, but already received offset 50 before", terr.ErrorMessage) + }) +} + +func TestStreamAbandoning(t *testing.T) { + connFC := newConnectionFlowController( + 100, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + require.True(t, connFC.UpdateSendWindow(300)) + fc := newStreamFlowController( + 42, + connFC, + 60, + protocol.MaxByteCount, + 100, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + + require.NoError(t, fc.UpdateHighestReceived(50, true, monotime.Now())) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) + require.Zero(t, connFC.GetWindowUpdate(monotime.Now())) + + // Abandon the stream. + // This marks all bytes as having been consumed. + fc.Abandon() + require.Equal(t, protocol.ByteCount(150), connFC.GetWindowUpdate(monotime.Now())) +} + +func TestStreamSendWindow(t *testing.T) { + // We set up the connection flow controller with a limit of 300 bytes, + // and the stream flow controller with a limit of 100 bytes. + connFC := newConnectionFlowController( + protocol.MaxByteCount, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + require.True(t, connFC.UpdateSendWindow(300)) + fc := newStreamFlowController( + 42, + connFC, + protocol.MaxByteCount, + protocol.MaxByteCount, + 100, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + // first, we're limited by the stream flow controller + require.Equal(t, protocol.ByteCount(100), fc.SendWindowSize()) + require.True(t, fc.TryAddBytesSent(50)) + require.False(t, fc.IsNewlyBlocked()) + require.Equal(t, protocol.ByteCount(50), fc.SendWindowSize()) + require.True(t, fc.TryAddBytesSent(50)) + require.True(t, fc.IsNewlyBlocked()) + require.Zero(t, fc.SendWindowSize()) + require.False(t, fc.IsNewlyBlocked()) // we're still blocked, but it's not new + + // Update the stream flow control limit, but don't update the connection flow control limit. + // We're now limited by the connection flow controller. + require.True(t, fc.UpdateSendWindow(1000)) + // reordered updates are ignored + require.False(t, fc.UpdateSendWindow(999)) + + require.False(t, fc.IsNewlyBlocked()) // we're not blocked anymore + require.Equal(t, protocol.ByteCount(200), fc.SendWindowSize()) + require.True(t, fc.TryAddBytesSent(200)) + require.Zero(t, fc.SendWindowSize()) + require.False(t, fc.IsNewlyBlocked()) // we're blocked, but not on stream flow control +} + +func TestStreamFlowControllerTryAddBytesSent(t *testing.T) { + connFC := newConnectionFlowController( + protocol.MaxByteCount, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + require.True(t, connFC.UpdateSendWindow(10)) + fc := newStreamFlowController( + 42, + connFC, + protocol.MaxByteCount, + protocol.MaxByteCount, + 100, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + + require.True(t, fc.TryAddBytesSent(6)) + require.False(t, fc.TryAddBytesSent(5)) + require.Equal(t, protocol.ByteCount(4), connFC.SendWindowSize()) +} + +func TestStreamWindowUpdate(t *testing.T) { + fc := newStreamFlowController( + 42, + newConnectionFlowController( + protocol.MaxByteCount, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ), + 100, + 100, + protocol.MaxByteCount, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) + hasStreamWindowUpdate, _ := fc.AddBytesRead(24) + require.False(t, hasStreamWindowUpdate) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) + // the window is updated when it's 25% filled + hasStreamWindowUpdate, _ = fc.AddBytesRead(1) + require.True(t, hasStreamWindowUpdate) + require.Equal(t, protocol.ByteCount(125), fc.GetWindowUpdate(monotime.Now())) + + hasStreamWindowUpdate, _ = fc.AddBytesRead(24) + require.False(t, hasStreamWindowUpdate) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) + // the window is updated when it's 25% filled + hasStreamWindowUpdate, _ = fc.AddBytesRead(1) + require.True(t, hasStreamWindowUpdate) + require.Equal(t, protocol.ByteCount(150), fc.GetWindowUpdate(monotime.Now())) + + // Receive the final offset. + // We don't need to send any more flow control updates. + require.NoError(t, fc.UpdateHighestReceived(100, true, monotime.Now())) + fc.AddBytesRead(50) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) +} + +func TestStreamConnectionWindowUpdate(t *testing.T) { + connFC := newConnectionFlowController( + 100, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + fc := newStreamFlowController( + 42, + connFC, + 1000, + protocol.MaxByteCount, + protocol.MaxByteCount, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + + hasStreamWindowUpdate, hasConnWindowUpdate := fc.AddBytesRead(50) + require.False(t, hasStreamWindowUpdate) + require.Zero(t, fc.GetWindowUpdate(monotime.Now())) + require.True(t, hasConnWindowUpdate) + require.NotZero(t, connFC.GetWindowUpdate(monotime.Now())) +} + +func TestStreamWindowAutoTuning(t *testing.T) { + // the RTT is 1 second + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(time.Second, 0) + require.Equal(t, time.Second, rttStats.SmoothedRTT()) + + connFC := newConnectionFlowController( + 150, // initial receive window + 350, // max receive window + func(size protocol.ByteCount) bool { return true }, + rttStats, + utils.DefaultLogger, + ) + fc := newStreamFlowController( + 42, + connFC, + 100, // initial send window + 399, // max send window + protocol.MaxByteCount, + rttStats, + utils.DefaultLogger, + ) + + now := monotime.Now() + require.NoError(t, fc.UpdateHighestReceived(100, false, now)) + + // data consumption is too slow, window size is not increased + now = now.Add(2500 * time.Millisecond) + fc.AddBytesRead(51) + // one initial stream window size added + require.Equal(t, protocol.ByteCount(51+100), fc.GetWindowUpdate(now)) + // one initial connection window size added + require.Equal(t, protocol.ByteCount(51+150), connFC.getWindowUpdate(now)) + + // data consumption is fast enough, window size is increased + now = now.Add(2 * time.Second) + fc.AddBytesRead(51) + // stream window size doubled to 200 bytes + require.Equal(t, protocol.ByteCount(102+2*100), fc.GetWindowUpdate(now)) + // The connection window is now increased as well, + // so that we don't get blocked on connection level flow control: + // The increase is by 200 bytes * a connection factor of 1.5: 300 bytes. + require.Equal(t, protocol.ByteCount(102+300), connFC.GetWindowUpdate(now)) + + // data consumption is fast enough, window size is increased + now = now.Add(2 * time.Second) + fc.AddBytesRead(101) + // stream window size increased again, but bumps into its maximum value + require.Equal(t, protocol.ByteCount(203+399), fc.GetWindowUpdate(now)) + // the connection window is also increased, but it bumps into its maximum value + require.Equal(t, protocol.ByteCount(203+350), connFC.GetWindowUpdate(now)) +} diff --git a/third_party/quic-go/flow_controller_test_helpers_test.go b/third_party/quic-go/flow_controller_test_helpers_test.go new file mode 100644 index 0000000..01d87de --- /dev/null +++ b/third_party/quic-go/flow_controller_test_helpers_test.go @@ -0,0 +1,39 @@ +package quic + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +func newTestStreamFlowController(id protocol.StreamID) *streamFlowController { + return newTestStreamFlowControllerWithSendWindow(id, 0) +} + +func newTestStreamFlowControllerWithSendWindow(id protocol.StreamID, sendWindow protocol.ByteCount) *streamFlowController { + return newTestStreamFlowControllerWithWindows(id, sendWindow, protocol.MaxByteCount, protocol.MaxByteCount) +} + +func newTestStreamFlowControllerWithWindows( + id protocol.StreamID, + sendWindow protocol.ByteCount, + streamReceiveWindow protocol.ByteCount, + connReceiveWindow protocol.ByteCount, +) *streamFlowController { + connFC := newConnectionFlowController( + connReceiveWindow, + protocol.MaxByteCount, + nil, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + connFC.UpdateSendWindow(protocol.MaxByteCount) + return newStreamFlowController( + id, + connFC, + streamReceiveWindow, + protocol.MaxByteCount, + sendWindow, + utils.NewRTTStats(), + utils.DefaultLogger, + ) +} diff --git a/third_party/quic-go/frame_sorter.go b/third_party/quic-go/frame_sorter.go new file mode 100644 index 0000000..417cdaf --- /dev/null +++ b/third_party/quic-go/frame_sorter.go @@ -0,0 +1,256 @@ +package quic + +import ( + "errors" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/utils/tree" +) + +// byteInterval is an interval from one ByteCount to the other +type byteInterval struct { + Start protocol.ByteCount + End protocol.ByteCount +} + +type frameSorterEntry struct { + Data []byte + DoneCb func() +} + +type frameSorter struct { + queue map[protocol.ByteCount]frameSorterEntry + readPos protocol.ByteCount + gapTree *tree.Btree[utils.ByteInterval] +} + +var errDuplicateStreamData = errors.New("duplicate stream data") + +func newFrameSorter() *frameSorter { + s := frameSorter{ + gapTree: tree.New[utils.ByteInterval](), + queue: make(map[protocol.ByteCount]frameSorterEntry), + } + s.gapTree.Insert(utils.ByteInterval{Start: 0, End: protocol.MaxByteCount}) + return &s +} + +func (s *frameSorter) Push(data []byte, offset protocol.ByteCount, doneCb func()) error { + err := s.push(data, offset, doneCb) + if err == errDuplicateStreamData { + if doneCb != nil { + doneCb() + } + return nil + } + return err +} + +func (s *frameSorter) push(data []byte, offset protocol.ByteCount, doneCb func()) error { + if len(data) == 0 { + return errDuplicateStreamData + } + + start := offset + end := offset + protocol.ByteCount(len(data)) + covInterval := utils.ByteInterval{Start: start, End: end} + + gaps := s.gapTree.Match(covInterval) + + if len(gaps) == 0 { + // No overlap with any existing gap + return errDuplicateStreamData + } + + startGap := gaps[0] + endGap := gaps[len(gaps)-1] + startGapEqualsEndGap := len(gaps) == 1 + + if startGapEqualsEndGap && end <= startGap.Start { + return errDuplicateStreamData + } + + startsInGap := covInterval.Start >= startGap.Start && covInterval.Start <= startGap.End + endsInGap := covInterval.End >= endGap.Start && covInterval.End < endGap.End + + startGapEnd := startGap.End // save it, in case startGap is modified + endGapStart := endGap.Start // save it, in case endGap is modified + endGapEnd := endGap.End // save it, in case endGap is modified + + var adjustedStartGapEnd bool + var wasCut bool + + pos := start + var hasReplacedAtLeastOne bool + for { + oldEntry, ok := s.queue[pos] + if !ok { + break + } + oldEntryLen := protocol.ByteCount(len(oldEntry.Data)) + if end-pos > oldEntryLen || (hasReplacedAtLeastOne && end-pos == oldEntryLen) { + // The existing frame is shorter than the new frame. Replace it. + delete(s.queue, pos) + pos += oldEntryLen + hasReplacedAtLeastOne = true + if oldEntry.DoneCb != nil { + oldEntry.DoneCb() + } + } else { + if !hasReplacedAtLeastOne { + return errDuplicateStreamData + } + // The existing frame is longer than the new frame. + // Cut the new frame such that the end aligns with the start of the existing frame. + data = data[:pos-start] + end = pos + wasCut = true + break + } + } + + if !startsInGap && !hasReplacedAtLeastOne { + // cut the frame, such that it starts at the start of the gap + data = data[startGap.Start-start:] + start = startGap.Start + wasCut = true + } + if start <= startGap.Start { + if end >= startGap.End { + // The frame covers the whole startGap. Delete the gap. + s.gapTree.Delete(startGap) + } else { + s.gapTree.Delete(startGap) + startGap.Start = end + // Re-insert the gap, but with the new start. + s.gapTree.Insert(startGap) + } + } else if !hasReplacedAtLeastOne { + s.gapTree.Delete(startGap) + startGap.End = start + // Re-insert the gap, but with the new end. + s.gapTree.Insert(startGap) + adjustedStartGapEnd = true + } + + if !startGapEqualsEndGap { + s.deleteConsecutive(startGapEnd) + for _, g := range gaps[1:] { + if g.End >= endGapStart { + break + } + s.deleteConsecutive(g.End) + s.gapTree.Delete(g) + } + } + + if !endsInGap && start != endGapEnd && end > endGapEnd { + // cut the frame, such that it ends at the end of the gap + data = data[:endGapEnd-start] + end = endGapEnd + wasCut = true + } + if end == endGapEnd { + if !startGapEqualsEndGap { + // The frame covers the whole endGap. Delete the gap. + s.gapTree.Delete(endGap) + } + } else { + if startGapEqualsEndGap && adjustedStartGapEnd { + // The frame split the existing gap into two. + s.gapTree.Insert(utils.ByteInterval{Start: end, End: startGapEnd}) + } else if !startGapEqualsEndGap { + s.gapTree.Delete(endGap) + endGap.Start = end + // Re-insert the gap, but with the new start. + s.gapTree.Insert(endGap) + } + } + + if wasCut && len(data) < protocol.MinStreamFrameBufferSize { + newData := make([]byte, len(data)) + copy(newData, data) + data = newData + if doneCb != nil { + doneCb() + doneCb = nil + } + } + + if s.gapTree.Len() > protocol.MaxStreamFrameSorterGaps { + return errors.New("too many gaps in received data") + } + + s.queue[start] = frameSorterEntry{Data: data, DoneCb: doneCb} + return nil +} + +// deleteConsecutive deletes consecutive frames from the queue, starting at pos +func (s *frameSorter) deleteConsecutive(pos protocol.ByteCount) { + for { + oldEntry, ok := s.queue[pos] + if !ok { + break + } + oldEntryLen := protocol.ByteCount(len(oldEntry.Data)) + delete(s.queue, pos) + if oldEntry.DoneCb != nil { + oldEntry.DoneCb() + } + pos += oldEntryLen + } +} + +func (s *frameSorter) Pop() (protocol.ByteCount, []byte, func()) { + entry, ok := s.queue[s.readPos] + if !ok { + return s.readPos, nil, nil + } + delete(s.queue, s.readPos) + offset := s.readPos + s.readPos += protocol.ByteCount(len(entry.Data)) + return offset, entry.Data, entry.DoneCb +} + +// HasMoreData says if there is any more data queued at *any* offset. +func (s *frameSorter) HasMoreData() bool { + return len(s.queue) > 0 +} + +var errTooLittleData = errors.New("too little data") + +// Peek copies len(p) consecutive bytes starting at offset into p, without removing them. +// It is only possible to peek from an offset where a frame starts. +// +// If there isn't enough consecutive data available, errTooLittleData is returned. +func (s *frameSorter) Peek(offset protocol.ByteCount, p []byte) error { + if len(p) == 0 { + return nil + } + + // first, check if we have enough consecutive data available + pos := offset + remaining := len(p) + for remaining > 0 { + entry, ok := s.queue[pos] + if !ok { + return errTooLittleData + } + entryLen := len(entry.Data) + if remaining <= entryLen { + break // enough data available + } + remaining -= entryLen + pos += protocol.ByteCount(entryLen) + } + + pos = offset + var copied int + for copied < len(p) { + entry := s.queue[pos] // the entry is guaranteed to exist from the check above + copied += copy(p[copied:], entry.Data) + pos += protocol.ByteCount(len(entry.Data)) + } + return nil +} diff --git a/third_party/quic-go/frame_sorter_test.go b/third_party/quic-go/frame_sorter_test.go new file mode 100644 index 0000000..f3fdb24 --- /dev/null +++ b/third_party/quic-go/frame_sorter_test.go @@ -0,0 +1,1653 @@ +package quic + +import ( + rand "crypto/rand" + "math" + mrand "math/rand/v2" + "slices" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" + + "github.com/stretchr/testify/require" +) + +type callbackTracker struct { + called *bool + cb func() +} + +func (t *callbackTracker) WasCalled() bool { return *t.called } + +func getFrameSorterTestCallback(t *testing.T) (func(), callbackTracker) { + var called bool + cb := func() { + if called { + t.Fatal("double free") + } + called = true + } + return cb, callbackTracker{ + cb: cb, + called: &called, + } +} + +func TestFrameSorterSimpleCases(t *testing.T) { + s := newFrameSorter() + _, data, doneCb := s.Pop() + require.Nil(t, data) + require.Nil(t, doneCb) + + // empty frames are ignored + require.NoError(t, s.Push(nil, 0, nil)) + _, data, doneCb = s.Pop() + require.Nil(t, data) + require.Nil(t, doneCb) + + cb1, t1 := getFrameSorterTestCallback(t) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push([]byte("bar"), 3, cb2)) + require.True(t, s.HasMoreData()) + require.NoError(t, s.Push([]byte("foo"), 0, cb1)) + + offset, data, doneCb := s.Pop() + require.Equal(t, []byte("foo"), data) + require.Zero(t, offset) + require.NotNil(t, doneCb) + doneCb() + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, s.HasMoreData()) + + offset, data, doneCb = s.Pop() + require.Equal(t, []byte("bar"), data) + require.Equal(t, protocol.ByteCount(3), offset) + require.NotNil(t, doneCb) + doneCb() + require.True(t, t2.WasCalled()) + require.False(t, s.HasMoreData()) + + // now receive a duplicate + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push([]byte("foo"), 0, cb3)) + require.False(t, s.HasMoreData()) + require.True(t, t3.WasCalled()) + + // now receive a later frame that overlaps with the ones we already consumed + cb4, _ := getFrameSorterTestCallback(t) + require.NoError(t, s.Push([]byte("barbaz"), 3, cb4)) + require.True(t, s.HasMoreData()) + + offset, data, _ = s.Pop() + require.Equal(t, protocol.ByteCount(6), offset) + require.Equal(t, []byte("baz"), data) + require.False(t, s.HasMoreData()) +} + +// Usually, it's not a good idea to test the implementation details. +// However, we need to make sure that the frame sorter handles gaps correctly, +// in particular when overlapping stream data is received. +// This also includes returning buffers that are no longer needed. +func TestFrameSorterGapHandling(t *testing.T) { + random := mrand.NewChaCha8([32]byte{'f', 'o', 'o', 'b', 'a', 'r'}) + + getData := func(l protocol.ByteCount) []byte { + b := make([]byte, l) + random.Read(b) + return b + } + + checkQueue := func(t *testing.T, s *frameSorter, m map[protocol.ByteCount][]byte) { + require.Equal(t, len(m), len(s.queue)) + for offset, data := range m { + require.Contains(t, s.queue, offset) + require.Equal(t, data, s.queue[offset].Data) + } + } + + checkGaps := func(t *testing.T, s *frameSorter, expectedGaps []byteInterval) { + actualGaps := s.gapTree.Values() + require.Len(t, actualGaps, len(expectedGaps)) + for i, gap := range actualGaps { + require.Equal(t, expectedGaps[i], byteInterval{Start: gap.Start, End: gap.End}) + } + } + + // ---xxx-------------- + // ++++++ + // => + // ---xxx++++++-------- + t.Run("case 1", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 11, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ---xxx----------------- + // +++++++ + // => + // ---xxx---+++++++-------- + t.Run("case 2", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 10, cb2)) // 10 -15 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 10: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 6, End: 10}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ---xxx----xxxxxx------- + // ++++ + // => + // ---xxx++++xxxxx-------- + t.Run("case 3", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f3, 10, cb2)) // 10 - 15 + require.NoError(t, s.Push(f2, 6, cb3)) // 6 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f2, + 10: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ----xxxx------- + // ++++ + // => + // ----xxxx++----- + t.Run("case 4", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 7 + require.NoError(t, s.Push(f2, 5, cb2)) // 5 - 9 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 7: f2[2:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 9, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + t.Run("case 4, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 2)) + f1 := getData(4 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 7 + require.NoError(t, s.Push(f2, 5*mult, cb2)) // 5 - 9 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f1, + 7 * mult: f2[2*mult:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 9 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // xxxx------- + // ++++ + // => + // xxxx+++----- + t.Run("case 5", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 0, cb1)) // 0 - 4 + require.NoError(t, s.Push(f2, 3, cb2)) // 3 - 7 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 0: f1, + 4: f2[1:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 7, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + t.Run("case 5, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 2)) + f1 := getData(4 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 0, cb1)) // 0 - 4 + require.NoError(t, s.Push(f2, 3*mult, cb2)) // 3 - 7 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 0: f1, + 4 * mult: f2[mult:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 7 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ----xxxx------- + // ++++ + // => + // --++xxxx------- + t.Run("case 6", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 9 + require.NoError(t, s.Push(f2, 3, cb2)) // 3 - 7 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f2[:2], + 5: f1, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 9, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + t.Run("case 6, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 2)) + f1 := getData(4 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5*mult, cb1)) // 5 - 9 + require.NoError(t, s.Push(f2, 3*mult, cb2)) // 3 - 7 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f2[:2*mult], + 5 * mult: f1, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 9 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ---xxx----xxxxxx------- + // ++ + // => + // ---xxx++--xxxxx-------- + t.Run("case 7", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(2) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f3, 10, cb2)) // 10 - 15 + require.NoError(t, s.Push(f2, 6, cb3)) // 6 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f2, + 10: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 8, End: 10}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx---------xxxxxx-- + // ++ + // => + // ---xxx---++----xxxxx-- + t.Run("case 8", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(2) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f3, 15, cb2)) // 15 - 20 + require.NoError(t, s.Push(f2, 10, cb3)) // 10 - 12 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 10: f2, + 15: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 6, End: 10}, + {Start: 12, End: 15}, + {Start: 20, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx----xxxxxx------- + // ++ + // => + // ---xxx--++xxxxx-------- + t.Run("case 9", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(2) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f3, 10, cb2)) // 10 - 15 + require.NoError(t, s.Push(f2, 8, cb3)) // 8 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 8: f2, + 10: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 6, End: 8}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx----=====------- + // +++++++ + // => + // ---xxx++++=====-------- + t.Run("case 10", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 10, cb2)) // 10 - 15 + require.NoError(t, s.Push(f3, 5, cb3)) // 5 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f3[1:5], + 10: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 10, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 4)) + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 10*mult, cb2)) // 10 - 15 + require.NoError(t, s.Push(f3, 5*mult, cb3)) // 5 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f1, + 6 * mult: f3[mult : 5*mult], + 10 * mult: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 15 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxxx----=====------- + // ++++++ + // => + // ---xxx++++=====-------- + t.Run("case 11", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 7 + require.NoError(t, s.Push(f2, 10, cb2)) // 10 - 15 + require.NoError(t, s.Push(f3, 5, cb3)) // 5 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 7: f3[2:], + 10: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + // ---xxxx----=====------- + // ++++++ + // => + // ---xxx++++=====-------- + t.Run("case 11, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 3)) + f1 := getData(4 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 7 + require.NoError(t, s.Push(f2, 10*mult, cb2)) // 10 - 15 + require.NoError(t, s.Push(f3, 5*mult, cb3)) // 5 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f1, + 7 * mult: f3[2*mult:], + 10 * mult: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 15 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ----xxxx------- + // +++++++ + // => + // ----+++++++----- + t.Run("case 12", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(7) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 7 + require.NoError(t, s.Push(f2, 3, cb2)) // 3 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 10, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ----xxx===------- + // +++++++ + // => + // ----+++++++----- + t.Run("case 13", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(7) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 9 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 10, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ----xxx====------- + // +++++ + // => + // ----+++====----- + t.Run("case 14", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 10 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f3[:3], + 6: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 10, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 14, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 3)) + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6*mult, cb2)) // 6 - 10 + require.NoError(t, s.Push(f3, 3*mult, cb3)) // 3 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f3[:3*mult], + 6 * mult: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 10 * mult, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ----xxx===------- + // ++++++ + // => + // ----++++++----- + t.Run("case 15", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 9 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 9 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 9, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxxx------- + // ++++ + // => + // ---xxxx----- + t.Run("case 16", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 9 + require.NoError(t, s.Push(f2, 5, cb2)) // 5 - 9 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 5: f1, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 5}, + {Start: 9, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + // ----xxx===------- + // +++ + // => + // ----xxx===----- + t.Run("case 17", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(3) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 9 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 6 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 9, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + // ---xxxx------- + // ++ + // => + // ---xxxx----- + t.Run("case 18", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(4) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(2) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 9 + require.NoError(t, s.Push(f2, 5, cb2)) // 5 - 7 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 5: f1, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 5}, + {Start: 9, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + // ---xxxxx------ + // ++ + // => + // ---xxxxx---- + t.Run("case 19", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(5) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(2) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 10 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 5: f1, + }) + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 5: f1, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 5}, + {Start: 10, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + // xxxxx------ + // ++ + // => + // xxxxx------ + t.Run("case 20", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(10) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 0, cb1)) // 0 - 10 + require.NoError(t, s.Push(f2, 5, cb2)) // 5 - 9 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 0: f1, + }) + checkGaps(t, s, []byteInterval{ + {Start: 10, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + // ---xxxxx--- + // +++ + // => + // ---xxxxx--- + t.Run("case 21", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(5) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 10 + require.NoError(t, s.Push(f2, 7, cb2)) // 7 - 10 + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 5}, + {Start: 10, End: protocol.MaxByteCount}, + }) + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 5: f1, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + }) + + // ----xxx------ + // +++++ + // => + // --+++++---- + t.Run("case 22", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 8 + require.NoError(t, s.Push(f2, 3, cb2)) // 3 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 8, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ----xxx===------ + // ++++++++ + // => + // --++++++++---- + t.Run("case 23", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(8) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 8 + require.NoError(t, s.Push(f2, 8, cb2)) // 8 - 11 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 11, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // --xxx---===--- + // ++++++ + // => + // --xxx++++++---- + t.Run("case 24", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 6, cb3)) // 6 - 12 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 12, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // --xxx---===---### + // +++++++++ + // => + // --xxx+++++++++### + t.Run("case 25", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(3) + cb3, t3 := getFrameSorterTestCallback(t) + f4 := getData(9) + cb4, t4 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 15, cb3)) // 15 - 18 + require.NoError(t, s.Push(f4, 6, cb4)) // 6 - 15 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f4, + 15: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 18, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + require.False(t, t4.WasCalled()) + }) + + // ----xxx------ + // +++++++ + // => + // --+++++++--- + t.Run("case 26", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(10) + cb2, t2 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 8 + require.NoError(t, s.Push(f2, 3, cb2)) // 3 - 13 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 13, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + }) + + // ---xxx====--- + // ++++ + // => + // --+xxx====--- + t.Run("case 27", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(4) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 10 + require.NoError(t, s.Push(f3, 2, cb3)) // 2 - 6 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 2: f3[:1], + 3: f1, + 6: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 2}, + {Start: 10, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 27, for long frames", func(t *testing.T) { + s := newFrameSorter() + const mult = protocol.MinStreamFrameSize + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(4 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6*mult, cb2)) // 6 - 10 + require.NoError(t, s.Push(f3, 2*mult, cb3)) // 2 - 6 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 2 * mult: f3[:mult], + 3 * mult: f1, + 6 * mult: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 2 * mult}, + {Start: 10 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx====--- + // ++++++ + // => + // --+xxx====--- + t.Run("case 28", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 10 + require.NoError(t, s.Push(f3, 2, cb3)) // 2 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 2: f3[:1], + 3: f1, + 6: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 2}, + {Start: 10, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 28, for long frames", func(t *testing.T) { + s := newFrameSorter() + const mult = protocol.MinStreamFrameSize + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6*mult, cb2)) // 6 - 10 + require.NoError(t, s.Push(f3, 2*mult, cb3)) // 2 - 8 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 2 * mult: f3[:mult], + 3 * mult: f1, + 6 * mult: f2, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 2 * mult}, + {Start: 10 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx===----- + // +++++ + // => + // ---xxx+++++--- + t.Run("case 29", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(5) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 9 + require.NoError(t, s.Push(f3, 6, cb3)) // 6 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 11, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx===---- + // ++++++ + // => + // ---xxx===++-- + t.Run("case 30", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6, cb2)) // 6 - 9 + require.NoError(t, s.Push(f3, 5, cb3)) // 5 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f2, + 9: f3[4:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 11, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 30, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 2)) + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(6 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 6*mult, cb2)) // 6 - 9 + require.NoError(t, s.Push(f3, 5*mult, cb3)) // 5 - 11 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f1, + 6 * mult: f2, + 9 * mult: f3[4*mult:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 11 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx---===----- + // ++++++++++ + // => + // ---xxx++++++++--- + t.Run("case 31", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(10) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 5, cb3)) // 5 - 15 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f3[1:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 15, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 31, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 9)) + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(10 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9*mult, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 5*mult, cb3)) // 5 - 15 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f1, + 6 * mult: f3[mult:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 15 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx---===----- + // +++++++++ + // => + // ---+++++++++--- + t.Run("case 32", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(9) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 12 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f3, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 12, End: protocol.MaxByteCount}, + }) + require.True(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) + + // ---xxx---===###----- + // ++++++++++++ + // => + // ---xxx++++++++++--- + t.Run("case 33", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(3) + cb3, t3 := getFrameSorterTestCallback(t) + f4 := getData(12) + cb4, t4 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 9, cb3)) // 12 - 15 + require.NoError(t, s.Push(f4, 5, cb4)) // 5 - 17 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f1, + 6: f4[1:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 17, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + require.True(t, t4.WasCalled()) + }) + + t.Run("case 33, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 11)) + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(3 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(3 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + f4 := getData(12 * mult) + cb4, t4 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9*mult, cb2)) // 9 - 12 + require.NoError(t, s.Push(f3, 9*mult, cb3)) // 12 - 15 + require.NoError(t, s.Push(f4, 5*mult, cb4)) // 5 - 17 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f1, + 6 * mult: f4[mult:], + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 17 * mult, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + require.False(t, t4.WasCalled()) + }) + + // ---xxx===---### + // ++++++ + // => + // ---xxx++++++### + t.Run("case 34", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(5) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(5) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(10) + cb3, t3 := getFrameSorterTestCallback(t) + f4 := getData(5) + cb4, t4 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 5, cb1)) // 5 - 10 + require.NoError(t, s.Push(f2, 10, cb2)) // 10 - 15 + require.NoError(t, s.Push(f4, 20, cb3)) // 20 - 25 + require.NoError(t, s.Push(f3, 10, cb4)) // 10 - 20 + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 5: f1, + 10: f3, + 20: f4, + }) + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 5}, + {Start: 25, End: protocol.MaxByteCount}, + }) + require.False(t, t1.WasCalled()) + require.True(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + require.False(t, t4.WasCalled()) + }) + + // ---xxx---####--- + // ++++++++ + // => + // ---++++++####--- + t.Run("case 35", func(t *testing.T) { + s := newFrameSorter() + f1 := getData(3) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(8) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9, cb2)) // 9 - 13 + require.NoError(t, s.Push(f3, 3, cb3)) // 3 - 11 + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3}, + {Start: 13, End: protocol.MaxByteCount}, + }) + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3: f3[:6], + 9: f2, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.True(t, t3.WasCalled()) + }) + + t.Run("case 35, for long frames", func(t *testing.T) { + s := newFrameSorter() + mult := protocol.ByteCount(math.Ceil(float64(protocol.MinStreamFrameSize) / 6)) + f1 := getData(3 * mult) + cb1, t1 := getFrameSorterTestCallback(t) + f2 := getData(4 * mult) + cb2, t2 := getFrameSorterTestCallback(t) + f3 := getData(8 * mult) + cb3, t3 := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f1, 3*mult, cb1)) // 3 - 6 + require.NoError(t, s.Push(f2, 9*mult, cb2)) // 9 - 13 + require.NoError(t, s.Push(f3, 3*mult, cb3)) // 3 - 11 + checkGaps(t, s, []byteInterval{ + {Start: 0, End: 3 * mult}, + {Start: 13 * mult, End: protocol.MaxByteCount}, + }) + checkQueue(t, s, map[protocol.ByteCount][]byte{ + 3 * mult: f3[:6*mult], + 9 * mult: f2, + }) + require.True(t, t1.WasCalled()) + require.False(t, t2.WasCalled()) + require.False(t, t3.WasCalled()) + }) +} + +func TestFrameSorterTooManyGaps(t *testing.T) { + s := newFrameSorter() + for i := range protocol.MaxStreamFrameSorterGaps { + require.NoError(t, s.Push([]byte("foobar"), protocol.ByteCount(i*7), nil)) + } + require.Equal(t, protocol.MaxStreamFrameSorterGaps, s.gapTree.Len()) + err := s.Push([]byte("foobar"), protocol.ByteCount(protocol.MaxStreamFrameSorterGaps*7)+100, nil) + require.EqualError(t, err, "too many gaps in received data") +} + +func TestFrameSorterRandomized(t *testing.T) { + t.Run("short", func(t *testing.T) { + testFrameSorterRandomized(t, 25, false, false) + }) + t.Run("long", func(t *testing.T) { + testFrameSorterRandomized(t, 2*protocol.MinStreamFrameSize, false, false) + }) + t.Run("short, with duplicates", func(t *testing.T) { + testFrameSorterRandomized(t, 25, true, false) + }) + t.Run("long, with duplicates", func(t *testing.T) { + testFrameSorterRandomized(t, 2*protocol.MinStreamFrameSize, true, false) + }) + t.Run("short, with overlaps", func(t *testing.T) { + testFrameSorterRandomized(t, 25, false, true) + }) + t.Run("long, with overlaps", func(t *testing.T) { + testFrameSorterRandomized(t, 2*protocol.MinStreamFrameSize, false, true) + }) +} + +func testFrameSorterRandomized(t *testing.T, dataLen protocol.ByteCount, injectDuplicates, injectOverlaps bool) { + type frame struct { + offset protocol.ByteCount + data []byte + } + + const num = 1000 + + data := make([]byte, num*int(dataLen)) + var seed [32]byte + rand.Read(seed[:]) + random := mrand.NewChaCha8(seed) + random.Read(data) + + frames := make([]frame, num) + for i := range num { + b := make([]byte, dataLen) + offset := i * int(dataLen) + copy(b, data[offset:offset+int(dataLen)]) + frames[i] = frame{ + offset: protocol.ByteCount(i) * dataLen, + data: b, + } + } + mrand.Shuffle(len(frames), func(i, j int) { frames[i], frames[j] = frames[j], frames[i] }) + + s := newFrameSorter() + + var callbacks []callbackTracker + for _, f := range frames { + cb, tr := getFrameSorterTestCallback(t) + require.NoError(t, s.Push(f.data, f.offset, cb)) + callbacks = append(callbacks, tr) + } + if injectDuplicates { + for range num / 10 { + cb, tr := getFrameSorterTestCallback(t) + df := frames[mrand.IntN(len(frames))] + require.NoError(t, s.Push(df.data, df.offset, cb)) + callbacks = append(callbacks, tr) + } + } + if injectOverlaps { + finalOffset := num * dataLen + for range num / 3 { + cb, tr := getFrameSorterTestCallback(t) + startOffset := protocol.ByteCount(mrand.IntN(int(finalOffset))) + endOffset := startOffset + protocol.ByteCount(mrand.IntN(int(finalOffset-startOffset))) + require.NoError(t, s.Push(data[startOffset:endOffset], startOffset, cb)) + callbacks = append(callbacks, tr) + } + } + require.Equal(t, 1, s.gapTree.Len()) + require.Equal(t, byteInterval{Start: num * dataLen, End: protocol.MaxByteCount}, func() byteInterval { + gap := s.gapTree.Head() + require.NotNil(t, gap) + return byteInterval{Start: gap.Start, End: gap.End} + }()) + + // read all data + var read []byte + for { + offset, b, cb := s.Pop() + if b == nil { + break + } + require.Equal(t, offset, protocol.ByteCount(len(read))) + read = append(read, b...) + if cb != nil { + cb() + } + } + + require.Equal(t, data, read) + require.False(t, s.HasMoreData()) + for _, cb := range callbacks { + require.True(t, cb.WasCalled()) + } +} + +func TestFrameSorterPeek(t *testing.T) { + s := newFrameSorter() + require.NoError(t, s.Peek(1337, []byte{})) // empty peek is a no-op + require.ErrorIs(t, s.Peek(0, []byte{0, 1, 2, 3, 4}), errTooLittleData) + + require.NoError(t, s.Push([]byte("foobar"), 0, nil)) + + // peek partial frame + p := make([]byte, 3) + require.NoError(t, s.Peek(0, p)) + require.Equal(t, []byte("foo"), p) + // peek entire frame + p = make([]byte, 6) + require.NoError(t, s.Peek(0, p)) + require.Equal(t, []byte("foobar"), p) + // peek more than available + p = make([]byte, 10) + require.ErrorIs(t, s.Peek(0, p), errTooLittleData) + // peek at offset where no entry exists + p = make([]byte, 3) + require.ErrorIs(t, s.Peek(3, p), errTooLittleData) + + // peek across multiple frames + s.Push([]byte("baz"), 6, nil) + p = make([]byte, 9) + require.NoError(t, s.Peek(0, p)) + require.Equal(t, []byte("foobarbaz"), p) + // peek starting from second frame + p = make([]byte, 3) + require.NoError(t, s.Peek(6, p)) + require.Equal(t, []byte("baz"), p) + + // peeking across gaps doesn't work + s.Push([]byte("qux"), 10, nil) + p = make([]byte, 10) + require.ErrorIs(t, s.Peek(0, p), errTooLittleData) +} + +func FuzzFrameSorter(f *testing.F) { + const ( + opPush uint8 = iota + opPop + opPeek + ) + + // Each operation is encoded as 3 bytes: 1 byte op type, 1 byte offset, 1 byte length. + // maxStreamLen is chosen larger than protocol.MinStreamFrameBufferSize so that overlapping + // frames can be cut down to less than that threshold. This ensures the small-cut code path + // in frameSorter.push (which copies the cut frame and fires the doneCb early) is exercised, + // while still allowing un-cut frames to be at or above the minimum size. + const ( + opSize = 3 + maxStreamLen = 256 + maxOps = 128 + ) + // If MinStreamFrameBufferSize is ever changed to grow beyond maxStreamLen, + // the small-cut path becomes unreachable. Keep them in sync. + require.Less(f, protocol.MinStreamFrameBufferSize, maxStreamLen) + + corpus := ossfuzzseeds.New(f) + for _, seed := range [][]uint8{ + {opPush, 0, 3, opPush, 3, 3, opPop, 0, 0}, + {opPush, 6, 3, opPush, 0, 3, opPush, 3, 3, opPop, 0, 0}, + {opPush, 0, 6, opPush, 0, 6, opPop, 0, 0, opPush, 0, 6}, + {opPush, 3, 4, opPush, 5, 4, opPush, 0, 3, opPop, 0, 0}, + {opPush, 3, 3, opPush, 9, 3, opPush, 5, 10, opPush, 0, 3, opPop, 0, 0}, + {opPush, 0, 6, opPeek, 0, 3, opPush, 6, 3, opPeek, 0, 9, opPush, 10, 3, opPeek, 0, 12}, + } { + corpus.Add(seed) + } + + // use deterministic non-uniform data + streamData := make([]byte, maxStreamLen) + for i := range streamData { + streamData[i] = byte(31*i + 7) + } + + f.Fuzz(func(t *testing.T, data []byte) { + if len(data)%opSize != 0 || len(data) > opSize*maxOps { + return + } + + s := newFrameSorter() + received := make([]bool, len(streamData)) + var readPos protocol.ByteCount + var callbacks []callbackTracker + + push := func(offset, length int) { + cb, tr := getFrameSorterTestCallback(t) + callbacks = append(callbacks, tr) + require.NoError(t, s.Push(streamData[offset:offset+length], protocol.ByteCount(offset), cb)) + for i := max(offset, int(readPos)); i < offset+length; i++ { + received[i] = true + } + } + + for ; len(data) >= opSize; data = data[opSize:] { + op, offset, length := data[0], int(data[1]), int(data[2]) + if offset+length > len(streamData) { + return + } + + switch op { + case opPush: + push(offset, length) + case opPop: + readPos = frameSorterFuzzPop(t, s, streamData, received, readPos) + case opPeek: + frameSorterFuzzPeek(t, s, streamData, received, readPos, offset, length) + } + } + + // Complete the stream so that all queued data is eventually popped or replaced. + // This lets us assert that every callback is called exactly once. + push(0, len(streamData)) + + for readPos < protocol.ByteCount(len(streamData)) { + readPos = frameSorterFuzzPop(t, s, streamData, received, readPos) + } + require.False(t, s.HasMoreData()) + for _, cb := range callbacks { + require.True(t, cb.WasCalled()) + } + }) +} + +func frameSorterFuzzPop(t *testing.T, s *frameSorter, streamData []byte, received []bool, readPos protocol.ByteCount) protocol.ByteCount { + t.Helper() + + hasMoreData := func(received []bool) bool { + return slices.Contains(received, true) + } + + offset, data, cb := s.Pop() + require.Equal(t, readPos, offset) + + require.LessOrEqual(t, readPos, protocol.ByteCount(len(streamData))) + if readPos == protocol.ByteCount(len(streamData)) || !received[readPos] { + require.Nil(t, data) + require.Nil(t, cb) + require.Equal(t, hasMoreData(received[int(readPos):]), s.HasMoreData()) + return readPos + } + + require.NotEmpty(t, data) + nextGap := slices.Index(received[int(readPos):], false) + if nextGap == -1 { + nextGap = len(received) - int(readPos) + } + require.LessOrEqual(t, len(data), nextGap) + end := readPos + protocol.ByteCount(len(data)) + require.LessOrEqual(t, end, protocol.ByteCount(len(streamData))) + require.Equal(t, streamData[readPos:end], data) + if cb != nil { + cb() + } + require.Equal(t, hasMoreData(received[int(end):]), s.HasMoreData()) + return end +} + +func frameSorterFuzzPeek(t *testing.T, s *frameSorter, streamData []byte, received []bool, readPos protocol.ByteCount, offset, length int) { + t.Helper() + + p := make([]byte, length) + err := s.Peek(protocol.ByteCount(offset), p) + + // Peek only supports peeking from offsets where a frame starts. The model only + // knows that frame boundaries are guaranteed at readPos, so when peeking from + // there with enough consecutive received data, Peek must succeed. + if length > 0 && protocol.ByteCount(offset) == readPos { + nextGap := slices.Index(received[offset:], false) + if nextGap == -1 || nextGap >= length { + require.NoError(t, err) + } + } + if err != nil { + return + } + require.Equal(t, streamData[offset:offset+length], p) +} diff --git a/third_party/quic-go/framer.go b/third_party/quic-go/framer.go new file mode 100644 index 0000000..9f434e6 --- /dev/null +++ b/third_party/quic-go/framer.go @@ -0,0 +1,295 @@ +package quic + +import ( + "slices" + "sync" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils/ringbuffer" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/quicvarint" +) + +const ( + maxPathResponses = 256 + maxControlFrames = 16 << 10 +) + +// This is the largest possible size of a stream-related control frame +// (which is the RESET_STREAM frame). +const maxStreamControlFrameSize = 25 + +type streamFrameGetter interface { + popStreamFrame(protocol.ByteCount, protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) +} + +type streamControlFrameGetter interface { + getControlFrame(monotime.Time) (_ ackhandler.Frame, ok, hasMore bool) +} + +type framer struct { + mutex sync.Mutex + + activeStreams map[protocol.StreamID]streamFrameGetter + streamQueue ringbuffer.RingBuffer[protocol.StreamID] + streamsWithControlFrames map[protocol.StreamID]streamControlFrameGetter + + controlFrameMutex sync.Mutex + controlFrames []wire.Frame + pathResponses []*wire.PathResponseFrame + connFlowController *connectionFlowController + queuedTooManyControlFrames bool +} + +func newFramer(connFlowController *connectionFlowController) *framer { + return &framer{ + activeStreams: make(map[protocol.StreamID]streamFrameGetter), + streamsWithControlFrames: make(map[protocol.StreamID]streamControlFrameGetter), + connFlowController: connFlowController, + } +} + +func (f *framer) HasData() bool { + f.mutex.Lock() + hasData := !f.streamQueue.Empty() + f.mutex.Unlock() + if hasData { + return true + } + f.controlFrameMutex.Lock() + defer f.controlFrameMutex.Unlock() + return len(f.streamsWithControlFrames) > 0 || len(f.controlFrames) > 0 || len(f.pathResponses) > 0 +} + +func (f *framer) QueueControlFrame(frame wire.Frame) { + f.controlFrameMutex.Lock() + defer f.controlFrameMutex.Unlock() + + if pr, ok := frame.(*wire.PathResponseFrame); ok { + // Only queue up to maxPathResponses PATH_RESPONSE frames. + // This limit should be high enough to never be hit in practice, + // unless the peer is doing something malicious. + if len(f.pathResponses) >= maxPathResponses { + return + } + f.pathResponses = append(f.pathResponses, pr) + return + } + // This is a hack. + if len(f.controlFrames) >= maxControlFrames { + f.queuedTooManyControlFrames = true + return + } + f.controlFrames = append(f.controlFrames, frame) +} + +func (f *framer) Append( + frames []ackhandler.Frame, + streamFrames []ackhandler.StreamFrame, + maxLen protocol.ByteCount, + now monotime.Time, + v protocol.Version, +) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) { + f.controlFrameMutex.Lock() + frames, controlFrameLen := f.appendControlFrames(frames, maxLen, now, v) + maxLen -= controlFrameLen + + var lastFrame ackhandler.StreamFrame + var streamFrameLen protocol.ByteCount + f.mutex.Lock() + // pop STREAM frames, until less than 128 bytes are left in the packet + numActiveStreams := f.streamQueue.Len() + for range numActiveStreams { + if protocol.MinStreamFrameSize > maxLen { + break + } + sf, blocked := f.getNextStreamFrame(maxLen, v) + if sf.Frame != nil { + streamFrames = append(streamFrames, sf) + maxLen -= sf.Frame.Length(v) + lastFrame = sf + streamFrameLen += sf.Frame.Length(v) + } + // If the stream just became blocked on stream flow control, attempt to pack the + // STREAM_DATA_BLOCKED into the same packet. + if blocked != nil { + l := blocked.Length(v) + // In case it doesn't fit, queue it for the next packet. + if maxLen < l { + f.controlFrames = append(f.controlFrames, blocked) + break + } + frames = append(frames, ackhandler.Frame{Frame: blocked}) + maxLen -= l + controlFrameLen += l + } + } + + // The only way to become blocked on connection-level flow control is by sending STREAM frames. + if isBlocked, offset := f.connFlowController.IsNewlyBlocked(); isBlocked { + blocked := &wire.DataBlockedFrame{MaximumData: offset} + l := blocked.Length(v) + // In case it doesn't fit, queue it for the next packet. + if maxLen >= l { + frames = append(frames, ackhandler.Frame{Frame: blocked}) + controlFrameLen += l + } else { + f.controlFrames = append(f.controlFrames, blocked) + } + } + + f.mutex.Unlock() + f.controlFrameMutex.Unlock() + + if lastFrame.Frame != nil { + // account for the smaller size of the last STREAM frame + streamFrameLen -= lastFrame.Frame.Length(v) + lastFrame.Frame.DataLenPresent = false + streamFrameLen += lastFrame.Frame.Length(v) + } + + return frames, streamFrames, controlFrameLen + streamFrameLen +} + +func (f *framer) appendControlFrames( + frames []ackhandler.Frame, + maxLen protocol.ByteCount, + now monotime.Time, + v protocol.Version, +) ([]ackhandler.Frame, protocol.ByteCount) { + var length protocol.ByteCount + // add a PATH_RESPONSE first, but only pack a single PATH_RESPONSE per packet + if len(f.pathResponses) > 0 { + frame := f.pathResponses[0] + frameLen := frame.Length(v) + if frameLen <= maxLen { + frames = append(frames, ackhandler.Frame{Frame: frame}) + length += frameLen + f.pathResponses = f.pathResponses[1:] + } + } + + // add stream-related control frames + for id, str := range f.streamsWithControlFrames { + start: + remainingLen := maxLen - length + if remainingLen <= maxStreamControlFrameSize { + break + } + fr, ok, hasMore := str.getControlFrame(now) + if !hasMore { + delete(f.streamsWithControlFrames, id) + } + if !ok { + continue + } + frames = append(frames, fr) + length += fr.Frame.Length(v) + if hasMore { + // It is rare that a stream has more than one control frame to queue. + // We don't want to spawn another loop for just to cover that case. + goto start + } + } + + for len(f.controlFrames) > 0 { + frame := f.controlFrames[len(f.controlFrames)-1] + frameLen := frame.Length(v) + if length+frameLen > maxLen { + break + } + frames = append(frames, ackhandler.Frame{Frame: frame}) + length += frameLen + f.controlFrames = f.controlFrames[:len(f.controlFrames)-1] + } + + return frames, length +} + +// QueuedTooManyControlFrames says if the control frame queue exceeded its maximum queue length. +// This is a hack. +// It is easier to implement than propagating an error return value in QueueControlFrame. +// The correct solution would be to queue frames with their respective structs. +// See https://github.com/apernet/quic-go/issues/4271 for the queueing of stream-related control frames. +func (f *framer) QueuedTooManyControlFrames() bool { + return f.queuedTooManyControlFrames +} + +func (f *framer) AddActiveStream(id protocol.StreamID, str streamFrameGetter) { + f.mutex.Lock() + if _, ok := f.activeStreams[id]; !ok { + f.streamQueue.PushBack(id) + f.activeStreams[id] = str + } + f.mutex.Unlock() +} + +func (f *framer) AddStreamWithControlFrames(id protocol.StreamID, str streamControlFrameGetter) { + f.controlFrameMutex.Lock() + if _, ok := f.streamsWithControlFrames[id]; !ok { + f.streamsWithControlFrames[id] = str + } + f.controlFrameMutex.Unlock() +} + +// RemoveActiveStream is called when a stream completes. +func (f *framer) RemoveActiveStream(id protocol.StreamID) { + f.mutex.Lock() + delete(f.activeStreams, id) + // We don't delete the stream from the streamQueue, + // since we'd have to iterate over the ringbuffer. + // Instead, we check if the stream is still in activeStreams when appending STREAM frames. + f.mutex.Unlock() +} + +func (f *framer) getNextStreamFrame(maxLen protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame) { + id := f.streamQueue.PopFront() + // This should never return an error. Better check it anyway. + // The stream will only be in the streamQueue, if it enqueued itself there. + str, ok := f.activeStreams[id] + // The stream might have been removed after being enqueued. + if !ok { + return ackhandler.StreamFrame{}, nil + } + // For the last STREAM frame, we'll remove the DataLen field later. + // Therefore, we can pretend to have more bytes available when popping + // the STREAM frame (which will always have the DataLen set). + maxLen += protocol.ByteCount(quicvarint.Len(uint64(maxLen))) + frame, blocked, hasMoreData := str.popStreamFrame(maxLen, v) + if hasMoreData { // put the stream back in the queue (at the end) + f.streamQueue.PushBack(id) + } else { // no more data to send. Stream is not active + delete(f.activeStreams, id) + } + // Note that the frame.Frame can be nil: + // * if the stream was canceled after it said it had data + // * the remaining size doesn't allow us to add another STREAM frame + return frame, blocked +} + +func (f *framer) Handle0RTTRejection() { + f.mutex.Lock() + defer f.mutex.Unlock() + f.controlFrameMutex.Lock() + defer f.controlFrameMutex.Unlock() + + f.streamQueue.Clear() + for id := range f.activeStreams { + delete(f.activeStreams, id) + } + clear(f.streamsWithControlFrames) + var j int + for i, frame := range f.controlFrames { + switch frame.(type) { + case *wire.MaxDataFrame, *wire.MaxStreamDataFrame, *wire.MaxStreamsFrame, + *wire.DataBlockedFrame, *wire.StreamDataBlockedFrame, *wire.StreamsBlockedFrame: + continue + default: + f.controlFrames[j] = f.controlFrames[i] + j++ + } + } + f.controlFrames = slices.Delete(f.controlFrames, j, len(f.controlFrames)) +} diff --git a/third_party/quic-go/framer_test.go b/third_party/quic-go/framer_test.go new file mode 100644 index 0000000..d69cee8 --- /dev/null +++ b/third_party/quic-go/framer_test.go @@ -0,0 +1,476 @@ +package quic + +import ( + "bytes" + "encoding/binary" + "math/rand/v2" + "testing" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func TestFramerControlFrames(t *testing.T) { + pc := &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 6, 7, 8}} + msf := &wire.MaxStreamsFrame{MaxStreamNum: 0x1337} + + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + require.False(t, framer.HasData()) + framer.QueueControlFrame(pc) + require.True(t, framer.HasData()) + framer.QueueControlFrame(msf) + frames, streamFrames, length := framer.Append( + []ackhandler.Frame{{Frame: &wire.PingFrame{}}}, + nil, + protocol.MaxByteCount, + monotime.Now(), + protocol.Version1, + ) + require.Len(t, frames, 3) + require.Empty(t, streamFrames) + require.Contains(t, frames, ackhandler.Frame{Frame: &wire.PingFrame{}}) + require.Contains(t, frames, ackhandler.Frame{Frame: pc}) + require.Contains(t, frames, ackhandler.Frame{Frame: msf}) + require.Equal(t, length, pc.Length(protocol.Version1)+msf.Length(protocol.Version1)) + require.False(t, framer.HasData()) +} + +func TestFramerControlFrameSizing(t *testing.T) { + const maxSize = protocol.ByteCount(1000) + bf := &wire.DataBlockedFrame{MaximumData: 0x1337} + bfLen := bf.Length(protocol.Version1) + + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + numFrames := int(maxSize / bfLen) // max number of frames that fit into maxSize + for i := 0; i < numFrames+1; i++ { + framer.QueueControlFrame(bf) + } + frames, _, length := framer.Append(nil, nil, maxSize, monotime.Now(), protocol.Version1) + require.Len(t, frames, numFrames) + require.Greater(t, length, maxSize-bfLen) + // now make sure that the last frame is also added + frames, _, length = framer.Append(nil, nil, maxSize, monotime.Now(), protocol.Version1) + require.Len(t, frames, 1) + require.Equal(t, length, bfLen) +} + +func TestFramerStreamControlFrames(t *testing.T) { + const streamID = protocol.StreamID(10) + ping := &wire.PingFrame{} + mdf1 := &wire.MaxStreamDataFrame{StreamID: streamID, MaximumStreamData: 1337} + mdf2 := &wire.MaxStreamDataFrame{StreamID: streamID, MaximumStreamData: 1338} + + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + framer.QueueControlFrame(ping) + str := NewMockStreamControlFrameGetter(gomock.NewController(t)) + framer.AddStreamWithControlFrames(streamID, str) + now := monotime.Now() + str.EXPECT().getControlFrame(now).Return(ackhandler.Frame{Frame: mdf1}, true, true) + str.EXPECT().getControlFrame(now).Return(ackhandler.Frame{Frame: mdf2}, true, false) + frames, streamFrames, l := framer.Append(nil, nil, protocol.MaxByteCount, now, protocol.Version1) + require.Len(t, frames, 3) + require.Empty(t, streamFrames) + require.Equal(t, mdf1, frames[0].Frame) + require.Equal(t, mdf2, frames[1].Frame) + require.Equal(t, ping, frames[2].Frame) + require.Equal(t, ping.Length(protocol.Version1)+mdf1.Length(protocol.Version1)+mdf2.Length(protocol.Version1), l) +} + +// If there are less than 25 bytes left, no more stream-related control frames are enqueued. +// This avoids dequeueing a frame from the stream that would be too large to fit into the packet. +func TestFramerStreamControlFramesSizing(t *testing.T) { + mdf1 := &wire.MaxStreamDataFrame{MaximumStreamData: 1337} + + str := NewMockStreamControlFrameGetter(gomock.NewController(t)) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + framer.AddStreamWithControlFrames(10, str) + str.EXPECT().getControlFrame(gomock.Any()).Return(ackhandler.Frame{Frame: mdf1}, true, true).AnyTimes() + frames, _, l := framer.Append(nil, nil, 100, monotime.Now(), protocol.Version1) + require.Equal(t, protocol.ByteCount(len(frames))*mdf1.Length(protocol.Version1), l) + require.Greater(t, l, protocol.ByteCount(100-maxStreamControlFrameSize)) + require.LessOrEqual(t, l, protocol.ByteCount(100)) +} + +func TestFramerStreamDataBlocked(t *testing.T) { + t.Run("small STREAM frame", func(t *testing.T) { + testFramerStreamDataBlocked(t, true) + }) + + t.Run("large STREAM frame", func(t *testing.T) { + testFramerStreamDataBlocked(t, false) + }) +} + +// If the stream becomes blocked on stream flow control, we attempt to pack the STREAM_DATA_BLOCKED +// into the same packet. +// However, there's the pathological case, where the STREAM frame and the STREAM_DATA_BLOCKED frame +// don't fit into the same packet. In that case, the STREAM_DATA_BLOCKED frame is queued and sent +// in the next packet. +func testFramerStreamDataBlocked(t *testing.T, fits bool) { + const streamID = 5 + str := NewMockStreamFrameGetter(gomock.NewController(t)) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + framer.AddActiveStream(streamID, str) + str.EXPECT().popStreamFrame(gomock.Any(), gomock.Any()).DoAndReturn( + func(size protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + data := []byte("foobar") + if !fits { + // Leave 3 bytes in the packet. + // This is not enough to fit in the STREAM_DATA_BLOCKED frame. + data = make([]byte, size-3) + } + f := &wire.StreamFrame{StreamID: streamID, DataLenPresent: true, Data: data} + blocked := &wire.StreamDataBlockedFrame{StreamID: streamID, MaximumStreamData: f.DataLen()} + if !fits { + require.Greater(t, blocked.Length(protocol.Version1), protocol.ByteCount(3)) + } + return ackhandler.StreamFrame{Frame: f}, blocked, false + }, + ) + + const maxSize protocol.ByteCount = 1000 + frames, streamFrames, l := framer.Append(nil, nil, maxSize, monotime.Now(), protocol.Version1) + require.Len(t, streamFrames, 1) + dataLen := streamFrames[0].Frame.DataLen() + if fits { + require.Len(t, frames, 1) + require.Equal(t, &wire.StreamDataBlockedFrame{StreamID: streamID, MaximumStreamData: dataLen}, frames[0].Frame) + } else { + require.Equal(t, streamFrames[0].Frame.Length(protocol.Version1), l) + require.Empty(t, frames) + frames, streamFrames, l2 := framer.Append(nil, nil, maxSize, monotime.Now(), protocol.Version1) + require.Greater(t, l+l2, maxSize) + require.Empty(t, streamFrames) + require.Len(t, frames, 1) + require.Equal(t, &wire.StreamDataBlockedFrame{StreamID: streamID, MaximumStreamData: dataLen}, frames[0].Frame) + } +} + +func TestFramerDataBlocked(t *testing.T) { + t.Run("small STREAM frame", func(t *testing.T) { + testFramerDataBlocked(t, true) + }) + + t.Run("large STREAM frame", func(t *testing.T) { + testFramerDataBlocked(t, false) + }) +} + +// If the stream becomes blocked on connection flow control, we attempt to pack the +// DATA_BLOCKED frame into the same packet. +// However, there's the pathological case, where the STREAM frame and the DATA_BLOCKED frame +// don't fit into the same packet. In that case, the DATA_BLOCKED frame is queued and sent +// in the next packet. +func testFramerDataBlocked(t *testing.T, fits bool) { + const streamID = 5 + const offset = 100 + + fc := newConnectionFlowController(0, 0, nil, nil, nil) + fc.UpdateSendWindow(offset) + require.True(t, fc.TryAddBytesSent(offset)) + + str := NewMockStreamFrameGetter(gomock.NewController(t)) + framer := newFramer(fc) + framer.AddActiveStream(streamID, str) + + str.EXPECT().popStreamFrame(gomock.Any(), gomock.Any()).DoAndReturn( + func(size protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + data := []byte("foobar") + if !fits { + // Leave 2 bytes in the packet. + // This is not enough to fit in the DATA_BLOCKED frame. + data = make([]byte, size-2) + } + f := &wire.StreamFrame{StreamID: streamID, DataLenPresent: true, Data: data} + return ackhandler.StreamFrame{Frame: f}, nil, false + }, + ) + + const maxSize protocol.ByteCount = 1000 + frames, streamFrames, l := framer.Append(nil, nil, maxSize, monotime.Now(), protocol.Version1) + require.Len(t, streamFrames, 1) + if fits { + require.Len(t, frames, 1) + require.Equal(t, &wire.DataBlockedFrame{MaximumData: offset}, frames[0].Frame) + } else { + require.Equal(t, streamFrames[0].Frame.Length(protocol.Version1), l) + require.Empty(t, frames) + frames, streamFrames, l2 := framer.Append(nil, nil, maxSize, monotime.Now(), protocol.Version1) + require.Greater(t, l+l2, maxSize) + require.Empty(t, streamFrames) + require.Len(t, frames, 1) + require.Equal(t, &wire.DataBlockedFrame{MaximumData: offset}, frames[0].Frame) + } +} + +func TestFramerDetectsFrameDoS(t *testing.T) { + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + for i := range maxControlFrames - 1 { + framer.QueueControlFrame(&wire.PingFrame{}) + framer.QueueControlFrame(&wire.PingFrame{}) + require.False(t, framer.QueuedTooManyControlFrames()) + frames, _, _ := framer.Append([]ackhandler.Frame{}, nil, 1, monotime.Now(), protocol.Version1) + require.Len(t, frames, 1) + require.Len(t, framer.controlFrames, i+1) + } + framer.QueueControlFrame(&wire.PingFrame{}) + require.False(t, framer.QueuedTooManyControlFrames()) + require.Len(t, framer.controlFrames, maxControlFrames) + framer.QueueControlFrame(&wire.PingFrame{}) + require.True(t, framer.QueuedTooManyControlFrames()) + require.Len(t, framer.controlFrames, maxControlFrames) +} + +func TestFramerDetectsFramePathResponseDoS(t *testing.T) { + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + var pathResponses []*wire.PathResponseFrame + for range 2 * maxPathResponses { + var f wire.PathResponseFrame + binary.BigEndian.PutUint64(f.Data[:], rand.Uint64()) + pathResponses = append(pathResponses, &f) + framer.QueueControlFrame(&f) + } + for i := range maxPathResponses { + require.True(t, framer.HasData()) + frames, _, length := framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Len(t, frames, 1) + require.Equal(t, pathResponses[i], frames[0].Frame) + require.Equal(t, pathResponses[i].Length(protocol.Version1), length) + } + require.False(t, framer.HasData()) + frames, _, length := framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Empty(t, frames) + require.Zero(t, length) +} + +func TestFramerPacksSinglePathResponsePerPacket(t *testing.T) { + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + f1 := &wire.PathResponseFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}} + f2 := &wire.PathResponseFrame{Data: [8]byte{2, 3, 4, 5, 6, 7, 8, 9}} + cf1 := &wire.DataBlockedFrame{MaximumData: 1337} + cf2 := &wire.HandshakeDoneFrame{} + framer.QueueControlFrame(f1) + framer.QueueControlFrame(f2) + framer.QueueControlFrame(cf1) + framer.QueueControlFrame(cf2) + // the first packet should contain a single PATH_RESPONSE frame, but all the other control frames + frames, _, _ := framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Len(t, frames, 3) + require.Equal(t, f1, frames[0].Frame) + require.Contains(t, []wire.Frame{frames[1].Frame, frames[2].Frame}, cf1) + require.Contains(t, []wire.Frame{frames[1].Frame, frames[2].Frame}, cf2) + // the second packet should contain the other PATH_RESPONSE frame + require.True(t, framer.HasData()) + frames, _, _ = framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Len(t, frames, 1) + require.Equal(t, f2, frames[0].Frame) + require.False(t, framer.HasData()) +} + +func TestFramerAppendStreamFrames(t *testing.T) { + const ( + str1ID = protocol.StreamID(42) + str2ID = protocol.StreamID(43) + ) + f1 := &wire.StreamFrame{StreamID: str1ID, Data: []byte("foo"), DataLenPresent: true} + f2 := &wire.StreamFrame{StreamID: str2ID, Data: []byte("bar"), DataLenPresent: true} + totalLen := f1.Length(protocol.Version1) + f2.Length(protocol.Version1) + + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + require.False(t, framer.HasData()) + // no frames added yet + controlFrames, fs, length := framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Empty(t, controlFrames) + require.Empty(t, fs) + require.Zero(t, length) + + // add two streams + mockCtrl := gomock.NewController(t) + str1 := NewMockStreamFrameGetter(mockCtrl) + str1.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).Return(ackhandler.StreamFrame{Frame: f1}, nil, true) + str2 := NewMockStreamFrameGetter(mockCtrl) + str2.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).Return(ackhandler.StreamFrame{Frame: f2}, nil, false) + framer.AddActiveStream(str1ID, str1) + framer.AddActiveStream(str1ID, str1) // duplicate calls are ok (they're no-ops) + framer.AddActiveStream(str2ID, str2) + require.True(t, framer.HasData()) + + // Even though the first stream claimed to have more data, + // we only dequeue a single STREAM frame per call of AppendStreamFrames. + f0 := ackhandler.StreamFrame{Frame: &wire.StreamFrame{StreamID: 9999}} + controlFrames, fs, length = framer.Append([]ackhandler.Frame{}, []ackhandler.StreamFrame{f0}, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Empty(t, controlFrames) + require.Len(t, fs, 3) + require.Equal(t, f0, fs[0]) + require.Equal(t, str1ID, fs[1].Frame.StreamID) + require.Equal(t, []byte("foo"), fs[1].Frame.Data) + // since two STREAM frames are sent, the DataLenPresent flag is set on the first frame + require.True(t, fs[1].Frame.DataLenPresent) + require.Equal(t, str2ID, fs[2].Frame.StreamID) + require.Equal(t, []byte("bar"), fs[2].Frame.Data) + // the last frame doesn't have the DataLenPresent flag set + require.False(t, fs[2].Frame.DataLenPresent) + require.Equal(t, fs[1].Frame.Length(protocol.Version1)+fs[2].Frame.Length(protocol.Version1), length) + require.Less(t, length, totalLen) // unsetting DataLenPresent on the last frame reduces the length + require.True(t, framer.HasData()) // the stream claimed to have more data... + + // ... but it actually doesn't + str1.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).Return(ackhandler.StreamFrame{}, nil, false) + _, fs, length = framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Empty(t, fs) + require.Zero(t, length) + require.False(t, framer.HasData()) +} + +func TestFramerRemoveActiveStream(t *testing.T) { + const id = protocol.StreamID(42) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + require.False(t, framer.HasData()) + framer.AddActiveStream(id, NewMockStreamFrameGetter(gomock.NewController(t))) + require.True(t, framer.HasData()) + framer.RemoveActiveStream(id) // no calls will be issued to the mock stream + // we can't assert on framer.HasData here, since it's not removed from the ringbuffer + _, frames, _ := framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Empty(t, frames) + require.False(t, framer.HasData()) +} + +func TestFramerMinStreamFrameSize(t *testing.T) { + const id = protocol.StreamID(42) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + str := NewMockStreamFrameGetter(gomock.NewController(t)) + framer.AddActiveStream(id, str) + + require.True(t, framer.HasData()) + // don't pop frames smaller than the minimum STREAM frame size + _, frames, _ := framer.Append(nil, nil, protocol.MinStreamFrameSize-1, monotime.Now(), protocol.Version1) + require.Empty(t, frames) + + // pop frames of the minimum size + str.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).DoAndReturn( + func(size protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + f := &wire.StreamFrame{StreamID: id, DataLenPresent: true} + f.Data = make([]byte, f.MaxDataLen(protocol.MinStreamFrameSize, v)) + return ackhandler.StreamFrame{Frame: f}, nil, false + }, + ) + _, frames, _ = framer.Append(nil, nil, protocol.MinStreamFrameSize, monotime.Now(), protocol.Version1) + require.Len(t, frames, 1) + // unsetting DataLenPresent on the last frame reduced the size slightly beyond the minimum size + require.Equal(t, protocol.MinStreamFrameSize-2, frames[0].Frame.Length(protocol.Version1)) +} + +func TestFramerMinStreamFrameSizeMultipleStreamFrames(t *testing.T) { + const id = protocol.StreamID(42) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + str := NewMockStreamFrameGetter(gomock.NewController(t)) + framer.AddActiveStream(id, str) + + // pop a frame such that the remaining size is one byte less than the minimum STREAM frame size + f := &wire.StreamFrame{ + StreamID: id, + Data: bytes.Repeat([]byte("f"), int(500-protocol.MinStreamFrameSize)), + DataLenPresent: true, + } + str.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).Return(ackhandler.StreamFrame{Frame: f}, nil, false) + framer.AddActiveStream(id, str) + _, fs, length := framer.Append(nil, nil, 500, monotime.Now(), protocol.Version1) + require.Len(t, fs, 1) + require.Equal(t, f, fs[0].Frame) + require.Equal(t, f.Length(protocol.Version1), length) +} + +func TestFramerFillPacketOneStream(t *testing.T) { + const id = protocol.StreamID(42) + str := NewMockStreamFrameGetter(gomock.NewController(t)) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + + for i := protocol.MinStreamFrameSize; i < 2000; i++ { + str.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).DoAndReturn( + func(size protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + f := &wire.StreamFrame{ + StreamID: id, + DataLenPresent: true, + } + f.Data = make([]byte, f.MaxDataLen(size, v)) + require.Equal(t, size, f.Length(protocol.Version1)) + return ackhandler.StreamFrame{Frame: f}, nil, false + }, + ) + framer.AddActiveStream(id, str) + _, frames, _ := framer.Append(nil, nil, i, monotime.Now(), protocol.Version1) + require.Len(t, frames, 1) + require.False(t, frames[0].Frame.DataLenPresent) + // make sure the entire space was filled up + require.Equal(t, i, frames[0].Frame.Length(protocol.Version1)) + } +} + +func TestFramerFillPacketMultipleStreams(t *testing.T) { + const ( + id1 = protocol.StreamID(1000) + id2 = protocol.StreamID(11) + ) + mockCtrl := gomock.NewController(t) + stream1 := NewMockStreamFrameGetter(mockCtrl) + stream2 := NewMockStreamFrameGetter(mockCtrl) + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + + for i := 2 * protocol.MinStreamFrameSize; i < 2000; i++ { + stream1.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).DoAndReturn( + func(size protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + f := &wire.StreamFrame{StreamID: id1, DataLenPresent: true} + f.Data = make([]byte, f.MaxDataLen(protocol.MinStreamFrameSize, v)) + return ackhandler.StreamFrame{Frame: f}, nil, false + }, + ) + stream2.EXPECT().popStreamFrame(gomock.Any(), protocol.Version1).DoAndReturn( + func(size protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + f := &wire.StreamFrame{StreamID: id2, DataLenPresent: true} + f.Data = make([]byte, f.MaxDataLen(size, v)) + require.Equal(t, size, f.Length(protocol.Version1)) + return ackhandler.StreamFrame{Frame: f}, nil, false + }, + ) + framer.AddActiveStream(id1, stream1) + framer.AddActiveStream(id2, stream2) + _, frames, _ := framer.Append(nil, nil, i, monotime.Now(), protocol.Version1) + require.Len(t, frames, 2) + require.True(t, frames[0].Frame.DataLenPresent) + require.False(t, frames[1].Frame.DataLenPresent) + require.Equal(t, i, frames[0].Frame.Length(protocol.Version1)+frames[1].Frame.Length(protocol.Version1)) + } +} + +func TestFramer0RTTRejection(t *testing.T) { + ncid := &wire.NewConnectionIDFrame{ + SequenceNumber: 10, + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + } + ping := &wire.PingFrame{} + pc := &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 6, 7, 8}} + + framer := newFramer(newConnectionFlowController(0, 0, nil, nil, nil)) + framer.QueueControlFrame(ncid) + framer.QueueControlFrame(&wire.DataBlockedFrame{MaximumData: 1337}) + framer.QueueControlFrame(&wire.StreamDataBlockedFrame{StreamID: 42, MaximumStreamData: 1337}) + framer.QueueControlFrame(ping) + framer.QueueControlFrame(&wire.StreamsBlockedFrame{StreamLimit: 13}) + framer.QueueControlFrame(pc) + + framer.AddActiveStream(10, NewMockStreamFrameGetter(gomock.NewController(t))) + framer.AddStreamWithControlFrames(10, NewMockStreamControlFrameGetter(gomock.NewController(t))) + + framer.Handle0RTTRejection() + controlFrames, streamFrames, _ := framer.Append(nil, nil, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.Empty(t, streamFrames) + require.Len(t, controlFrames, 3) + require.Contains(t, controlFrames, ackhandler.Frame{Frame: pc}) + require.Contains(t, controlFrames, ackhandler.Frame{Frame: ping}) + require.Contains(t, controlFrames, ackhandler.Frame{Frame: ncid}) +} diff --git a/third_party/quic-go/go.mod b/third_party/quic-go/go.mod new file mode 100644 index 0000000..9a72bc0 --- /dev/null +++ b/third_party/quic-go/go.mod @@ -0,0 +1,35 @@ +module github.com/apernet/quic-go + +go 1.25.0 + +require ( + github.com/quic-go/go-ossfuzz-seeds v0.1.0 + github.com/quic-go/qpack v0.6.0 + github.com/refraction-networking/utls v1.8.2 + github.com/stretchr/testify v1.11.1 + go.uber.org/mock v0.5.2 + golang.org/x/crypto v0.54.0 + golang.org/x/net v0.56.0 + golang.org/x/sync v0.22.0 + golang.org/x/sys v0.47.0 +) + +require ( + github.com/andybalholm/brotli v1.0.6 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/jordanlewis/gcassert v0.0.0-20250430164644-389ef753e22e // indirect + github.com/klauspost/compress v1.18.7 // indirect + github.com/kr/pretty v0.3.1 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/rogpeppe/go-internal v1.10.0 // indirect + golang.org/x/mod v0.37.0 // indirect + golang.org/x/text v0.40.0 // indirect + golang.org/x/tools v0.47.0 // indirect + gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) + +tool ( + github.com/jordanlewis/gcassert/cmd/gcassert + go.uber.org/mock/mockgen +) diff --git a/third_party/quic-go/go.sum b/third_party/quic-go/go.sum new file mode 100644 index 0000000..f44495d --- /dev/null +++ b/third_party/quic-go/go.sum @@ -0,0 +1,107 @@ +github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sxfOI= +github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig= +github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= +github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/jordanlewis/gcassert v0.0.0-20250430164644-389ef753e22e h1:a+PGEeXb+exwBS3NboqXHyxarD9kaboBbrSp+7GuBuc= +github.com/jordanlewis/gcassert v0.0.0-20250430164644-389ef753e22e/go.mod h1:ZybsQk6DWyN5t7An1MuPm1gtSZ1xDaTXS9ZjIOxvQrk= +github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= +github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= +github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0= +github.com/quic-go/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk= +github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= +github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= +github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEvV+S9iJ2IdQo= +github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= +github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs= +github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ= +github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= +go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko= +go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o= +golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= +golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= +golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc= +golang.org/x/crypto v0.18.0/go.mod h1:R0j02AL6hcrfOiy9T4ZYp/rcWeMxM3L6QYxlOuEG1mg= +golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= +golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= +golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= +golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= +golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= +golang.org/x/mod v0.14.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c= +golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= +golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= +golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= +golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= +golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= +golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= +golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg= +golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk= +golang.org/x/net v0.20.0/go.mod h1:z8BVo6PvndSri0LbOE3hAn0apkU+1YvI6E70E9jsnvY= +golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= +golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= +golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y= +golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.16.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= +golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= +golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k= +golang.org/x/term v0.8.0/go.mod h1:xPskH00ivmX89bAKVGSKKtLOWNx2+17Eiy94tnKShWo= +golang.org/x/term v0.12.0/go.mod h1:owVbMEjm3cBLCHdkQu9b1opXd4ETQWc3BhuQGKgXgvU= +golang.org/x/term v0.16.0/go.mod h1:yn7UURbUtPyrVJPGPq404EukNFxcm/foM+bV/bfcDsY= +golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= +golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= +golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8= +golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE= +golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= +golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= +golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= +golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= +golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= +golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58= +golang.org/x/tools v0.17.0/go.mod h1:xsh6VxdV005rRVaS6SSAf9oiAqljS7UZUacMZ8Bnsps= +golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= +golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= +golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/third_party/quic-go/http3/README.md b/third_party/quic-go/http3/README.md new file mode 100644 index 0000000..eb09243 --- /dev/null +++ b/third_party/quic-go/http3/README.md @@ -0,0 +1,9 @@ +# HTTP/3 + +[![Documentation](https://img.shields.io/badge/docs-quic--go.net-red?style=flat)](https://quic-go.net/docs/) +[![PkgGoDev](https://pkg.go.dev/badge/github.com/apernet/quic-go/http3)](https://pkg.go.dev/github.com/apernet/quic-go/http3) + +This package implements HTTP/3 ([RFC 9114](https://datatracker.ietf.org/doc/html/rfc9114)), including QPACK ([RFC 9204](https://datatracker.ietf.org/doc/html/rfc9204)) and HTTP Datagrams ([RFC 9297](https://datatracker.ietf.org/doc/html/rfc9297)). +It aims to provide feature parity with the standard library's HTTP/1.1 and HTTP/2 implementation. + +Detailed documentation can be found on [quic-go.net](https://quic-go.net/docs/). diff --git a/third_party/quic-go/http3/body.go b/third_party/quic-go/http3/body.go new file mode 100644 index 0000000..254fec6 --- /dev/null +++ b/third_party/quic-go/http3/body.go @@ -0,0 +1,137 @@ +package http3 + +import ( + "context" + "errors" + "io" + "sync" + + "github.com/apernet/quic-go" +) + +// Settingser allows waiting for and retrieving the peer's HTTP/3 settings. +type Settingser interface { + // ReceivedSettings returns a channel that is closed once the peer's SETTINGS frame was received. + // Settings can be obtained from the Settings method after the channel was closed. + ReceivedSettings() <-chan struct{} + // Settings returns the settings received on this connection. + // It is only valid to call this function after the channel returned by ReceivedSettings was closed. + Settings() *Settings +} + +var errTooMuchData = errors.New("peer sent too much data") + +// The body is used in the requestBody (for a http.Request) and the responseBody (for a http.Response). +type body struct { + str *Stream + + remainingContentLength int64 + violatedContentLength bool + hasContentLength bool +} + +func newBody(str *Stream, contentLength int64) *body { + b := &body{str: str} + if contentLength >= 0 { + b.hasContentLength = true + b.remainingContentLength = contentLength + } + return b +} + +func (r *body) StreamID() quic.StreamID { return r.str.StreamID() } + +func (r *body) checkContentLengthViolation() error { + if !r.hasContentLength { + return nil + } + if r.remainingContentLength < 0 || r.remainingContentLength == 0 && r.str.hasMoreData() { + if !r.violatedContentLength { + r.str.CancelRead(quic.StreamErrorCode(ErrCodeMessageError)) + r.str.CancelWrite(quic.StreamErrorCode(ErrCodeMessageError)) + r.violatedContentLength = true + } + return errTooMuchData + } + return nil +} + +func (r *body) Read(b []byte) (int, error) { + if err := r.checkContentLengthViolation(); err != nil { + return 0, err + } + if r.hasContentLength { + b = b[:min(int64(len(b)), r.remainingContentLength)] + } + n, err := r.str.Read(b) + r.remainingContentLength -= int64(n) + if err := r.checkContentLengthViolation(); err != nil { + return n, err + } + return n, maybeReplaceError(err) +} + +func (r *body) Close() error { + r.str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled)) + return nil +} + +type requestBody struct { + body + connCtx context.Context + rcvdSettings <-chan struct{} + getSettings func() *Settings +} + +var _ io.ReadCloser = &requestBody{} + +func newRequestBody(str *Stream, contentLength int64, connCtx context.Context, rcvdSettings <-chan struct{}, getSettings func() *Settings) *requestBody { + return &requestBody{ + body: *newBody(str, contentLength), + connCtx: connCtx, + rcvdSettings: rcvdSettings, + getSettings: getSettings, + } +} + +type hijackableBody struct { + body body + + // only set for the http.Response + // The channel is closed when the user is done with this response: + // either when Read() errors, or when Close() is called. + reqDone chan<- struct{} + reqDoneOnce sync.Once +} + +var _ io.ReadCloser = &hijackableBody{} + +func newResponseBody(str *Stream, contentLength int64, done chan<- struct{}) *hijackableBody { + return &hijackableBody{ + body: *newBody(str, contentLength), + reqDone: done, + } +} + +func (r *hijackableBody) Read(b []byte) (int, error) { + n, err := r.body.Read(b) + if err != nil { + r.requestDone() + } + return n, maybeReplaceError(err) +} + +func (r *hijackableBody) requestDone() { + if r.reqDone != nil { + r.reqDoneOnce.Do(func() { + close(r.reqDone) + }) + } +} + +func (r *hijackableBody) Close() error { + r.requestDone() + // If the EOF was read, CancelRead() is a no-op. + r.body.str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled)) + return nil +} diff --git a/third_party/quic-go/http3/body_test.go b/third_party/quic-go/http3/body_test.go new file mode 100644 index 0000000..46a2f2a --- /dev/null +++ b/third_party/quic-go/http3/body_test.go @@ -0,0 +1,140 @@ +package http3 + +import ( + "bytes" + "io" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func TestResponseBodyReading(t *testing.T) { + mockCtrl := gomock.NewController(t) + var buf bytes.Buffer + buf.Write(getDataFrame([]byte("foobar"))) + str := NewMockDatagramStream(mockCtrl) + str.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + str.EXPECT().Read(gomock.Any()).DoAndReturn(buf.Read).AnyTimes() + reqDone := make(chan struct{}) + rb := newResponseBody( + newStream(str, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), + -1, + reqDone, + ) + + data, err := io.ReadAll(rb) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), data) +} + +func TestResponseBodyReadError(t *testing.T) { + mockCtrl := gomock.NewController(t) + str := NewMockDatagramStream(mockCtrl) + str.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + str.EXPECT().Read(gomock.Any()).Return(0, assert.AnError).Times(2) + reqDone := make(chan struct{}) + rb := newResponseBody( + newStream(str, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), + -1, + reqDone, + ) + + _, err := rb.Read([]byte{0}) + require.ErrorIs(t, err, assert.AnError) + // repeated calls to Read should return the same error + _, err = rb.Read([]byte{0}) + require.ErrorIs(t, err, assert.AnError) + select { + case <-reqDone: + default: + t.Fatal("reqDone should be closed") + } +} + +func TestResponseBodyClose(t *testing.T) { + mockCtrl := gomock.NewController(t) + str := NewMockDatagramStream(mockCtrl) + str.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + str.EXPECT().CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled)).Times(2) + reqDone := make(chan struct{}) + rb := newResponseBody( + newStream(str, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), + -1, + reqDone, + ) + require.NoError(t, rb.Close()) + select { + case <-reqDone: + default: + t.Fatal("reqDone should be closed") + } + + // multiple calls to Close should be a no-op + require.NoError(t, rb.Close()) +} + +func TestResponseBodyConcurrentClose(t *testing.T) { + mockCtrl := gomock.NewController(t) + str := NewMockDatagramStream(mockCtrl) + str.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + str.EXPECT().CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled)).MaxTimes(3) + reqDone := make(chan struct{}) + rb := newResponseBody( + newStream(str, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), + -1, + reqDone, + ) + + for range 3 { + go rb.Close() + } + select { + case <-reqDone: + case <-time.After(time.Second): + t.Fatal("reqDone should be closed") + } +} + +func TestResponseBodyLengthLimiting(t *testing.T) { + t.Run("along frame boundary", func(t *testing.T) { + testResponseBodyLengthLimiting(t, true) + }) + + t.Run("in the middle of a frame", func(t *testing.T) { + testResponseBodyLengthLimiting(t, false) + }) +} + +func testResponseBodyLengthLimiting(t *testing.T, alongFrameBoundary bool) { + var buf bytes.Buffer + buf.Write(getDataFrame([]byte("foo"))) + buf.Write(getDataFrame([]byte("bar"))) + + l := int64(4) + if alongFrameBoundary { + l = 3 + } + mockCtrl := gomock.NewController(t) + str := NewMockDatagramStream(mockCtrl) + str.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + str.EXPECT().CancelRead(quic.StreamErrorCode(ErrCodeMessageError)) + str.EXPECT().CancelWrite(quic.StreamErrorCode(ErrCodeMessageError)) + str.EXPECT().Read(gomock.Any()).DoAndReturn(buf.Read).AnyTimes() + rb := newResponseBody( + newStream(str, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), + l, + make(chan struct{}), + ) + data, err := io.ReadAll(rb) + require.Equal(t, []byte("foobar")[:l], data) + require.ErrorIs(t, err, errTooMuchData) + // check that repeated calls to Read also return the right error + n, err := rb.Read([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, errTooMuchData) +} diff --git a/third_party/quic-go/http3/capsule.go b/third_party/quic-go/http3/capsule.go new file mode 100644 index 0000000..2cfda7c --- /dev/null +++ b/third_party/quic-go/http3/capsule.go @@ -0,0 +1,148 @@ +package http3 + +import ( + "errors" + "io" + + "github.com/apernet/quic-go/quicvarint" +) + +// CapsuleType is the type of the capsule +type CapsuleType uint64 + +// CapsuleProtocolHeader is the header value used to advertise support for the capsule protocol +const CapsuleProtocolHeader = "Capsule-Protocol" + +type noCopy struct{} + +func (*noCopy) Lock() {} +func (*noCopy) Unlock() {} + +// CapsuleParser parses a sequence of capsules. +// A capsule's contents must be fully consumed or discarded before calling Next again. +type CapsuleParser struct { + noCopy noCopy + + r quicvarint.Reader + + generation uint64 + remaining uint64 +} + +// NewCapsuleParser creates a parser that reads capsules from r. +func NewCapsuleParser(r io.Reader) *CapsuleParser { + return &CapsuleParser{r: quicvarint.NewReader(r)} +} + +var ( + errReaderInvalid = errors.New("http3: capsule reader is no longer valid") + errCapsuleNotConsumed = errors.New("http3: previous capsule was not fully consumed") +) + +// Next returns the type and contents of the next capsule. +// The previous capsule's contents must be fully consumed or discarded before calling Next. +func (p *CapsuleParser) Next() (CapsuleType, CapsuleReader, error) { + if p.remaining > 0 { + return 0, CapsuleReader{}, errCapsuleNotConsumed + } + + r := &countingByteReader{Reader: p.r} + ct, err := quicvarint.Read(r) + if err != nil { + // If an io.EOF is returned without consuming any bytes, return it unmodified. + // Otherwise, return an io.ErrUnexpectedEOF. + if err == io.EOF && r.NumRead > 0 { + return 0, CapsuleReader{}, io.ErrUnexpectedEOF + } + return 0, CapsuleReader{}, err + } + r.Reset() + l, err := quicvarint.Read(r) + if err != nil { + if err == io.EOF { + return 0, CapsuleReader{}, io.ErrUnexpectedEOF + } + return 0, CapsuleReader{}, err + } + + p.generation++ + p.remaining = l + return CapsuleType(ct), CapsuleReader{parser: p, generation: p.generation}, nil +} + +// CapsuleReader reads the contents of a capsule. +// It becomes invalid when the parser advances to the next capsule. +type CapsuleReader struct { + parser *CapsuleParser + generation uint64 +} + +var _ quicvarint.Reader = CapsuleReader{} + +// valid reports whether the reader still refers to the parser's current capsule. +func (r CapsuleReader) valid() bool { + return r.parser != nil && r.generation == r.parser.generation +} + +// Read reads from the capsule contents. +func (r CapsuleReader) Read(b []byte) (int, error) { + if !r.valid() { + return 0, errReaderInvalid + } + if r.parser.remaining == 0 { + return 0, io.EOF + } + if uint64(len(b)) > r.parser.remaining { + b = b[:r.parser.remaining] + } + n, err := r.parser.r.Read(b) + r.parser.remaining -= uint64(n) + if err == io.EOF && r.parser.remaining > 0 { + return n, io.ErrUnexpectedEOF + } + return n, err +} + +// ReadByte reads one byte from the capsule contents. +func (r CapsuleReader) ReadByte() (byte, error) { + if !r.valid() { + return 0, errReaderInvalid + } + if r.parser.remaining == 0 { + return 0, io.EOF + } + b, err := r.parser.r.ReadByte() + if err == io.EOF { + return 0, io.ErrUnexpectedEOF + } + if err == nil { + r.parser.remaining-- + } + return b, err +} + +// Remaining returns the number of bytes remaining in the capsule. +func (r CapsuleReader) Remaining() int64 { + if !r.valid() { + return 0 + } + return int64(r.parser.remaining) +} + +// Discard consumes the remaining capsule contents. +func (r CapsuleReader) Discard() error { + _, err := io.Copy(io.Discard, r) + return err +} + +// WriteCapsule writes a capsule +func WriteCapsule(w quicvarint.Writer, ct CapsuleType, value []byte) error { + b := make([]byte, 0, 16) + b = quicvarint.Append(b, uint64(ct)) + b = quicvarint.Append(b, uint64(len(value))) + if _, err := w.Write(b); err != nil { + return err + } + _, err := w.Write(value) + return err +} diff --git a/third_party/quic-go/http3/capsule_test.go b/third_party/quic-go/http3/capsule_test.go new file mode 100644 index 0000000..8b385d4 --- /dev/null +++ b/third_party/quic-go/http3/capsule_test.go @@ -0,0 +1,163 @@ +package http3 + +import ( + "bytes" + "io" + "testing" + + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +func TestCapsuleParsing(t *testing.T) { + b := quicvarint.Append(nil, 1337) + b = quicvarint.Append(b, 6) + b = append(b, []byte("foobar")...) + + p := NewCapsuleParser(bytes.NewReader(b)) + ct, r, err := p.Next() + require.NoError(t, err) + require.Equal(t, CapsuleType(1337), ct) + require.Equal(t, int64(6), r.Remaining()) + buf := make([]byte, 3) + n, err := r.Read(buf) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("foo"), buf) + require.Equal(t, int64(3), r.Remaining()) + data, err := io.ReadAll(r) // reads until EOF + require.NoError(t, err) + require.Equal(t, []byte("bar"), data) + require.Zero(t, r.Remaining()) + _, _, err = p.Next() + require.ErrorIs(t, err, io.EOF) +} + +func TestEmptyCapsuleParsing(t *testing.T) { + b := quicvarint.Append(nil, 1337) + b = quicvarint.Append(b, 0) + // Capsule content is empty. + + p := NewCapsuleParser(bytes.NewReader(b)) + ct, r, err := p.Next() + require.NoError(t, err) + require.Equal(t, CapsuleType(1337), ct) + data, err := io.ReadAll(r) // reads until EOF + require.NoError(t, err) + require.Equal(t, []byte{}, data) +} + +// test EOF vs ErrUnexpectedEOF +func TestCapsuleTruncation(t *testing.T) { + t.Run("with content", func(t *testing.T) { + b := quicvarint.Append(nil, 1337) + b = quicvarint.Append(b, 6) + b = append(b, []byte("foobar")...) + testCapsuleTruncation(t, b) + }) + + t.Run("empty content", func(t *testing.T) { + b := quicvarint.Append(nil, 1337) + b = quicvarint.Append(b, 0) + testCapsuleTruncation(t, b) + }) +} + +func testCapsuleTruncation(t *testing.T, b []byte) { + for i := range b { + p := NewCapsuleParser(bytes.NewReader(b[:i])) + ct, r, err := p.Next() + if err != nil { + if i == 0 { + require.ErrorIs(t, err, io.EOF) + } else { + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + } + continue + } + require.Equal(t, CapsuleType(1337), ct) + _, err = io.ReadAll(r) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + } +} + +func TestCapsuleParserRequiresConsumption(t *testing.T) { + var buf bytes.Buffer + require.NoError(t, WriteCapsule(&buf, 1, []byte("first"))) + require.NoError(t, WriteCapsule(&buf, 2, []byte("second"))) + + p := NewCapsuleParser(&buf) + _, r, err := p.Next() + require.NoError(t, err) + _, err = r.ReadByte() + require.NoError(t, err) + _, _, err = p.Next() + require.ErrorIs(t, err, errCapsuleNotConsumed) + + require.NoError(t, r.Discard()) + ct, next, err := p.Next() + require.NoError(t, err) + require.Equal(t, CapsuleType(2), ct) + data, err := io.ReadAll(next) + require.NoError(t, err) + require.Equal(t, []byte("second"), data) + + _, err = r.ReadByte() + require.ErrorIs(t, err, errReaderInvalid) + _, err = r.Read(make([]byte, 1)) + require.ErrorIs(t, err, errReaderInvalid) +} + +func TestCopiedCapsuleReadersShareProgress(t *testing.T) { + var buf bytes.Buffer + require.NoError(t, WriteCapsule(&buf, 1, []byte("foobar"))) + + p := NewCapsuleParser(&buf) + _, r, err := p.Next() + require.NoError(t, err) + r2 := r + + b, err := r.ReadByte() + require.NoError(t, err) + require.Equal(t, byte('f'), b) + require.Equal(t, int64(5), r2.Remaining()) + data, err := io.ReadAll(r2) + require.NoError(t, err) + require.Equal(t, []byte("oobar"), data) + require.Zero(t, r.Remaining()) +} + +func TestCapsuleWriting(t *testing.T) { + var buf bytes.Buffer + require.NoError(t, WriteCapsule(&buf, 1337, []byte("foobar"))) + + p := NewCapsuleParser(&buf) + ct, r, err := p.Next() + require.NoError(t, err) + require.Equal(t, CapsuleType(1337), ct) + val, err := io.ReadAll(r) + require.NoError(t, err) + require.Equal(t, "foobar", string(val)) +} + +func TestCapsuleWriteEmpty(t *testing.T) { + var buf bytes.Buffer + require.NoError(t, WriteCapsule(&buf, 1337, []byte{})) + require.NoError(t, WriteCapsule(&buf, 1337, []byte{})) + + p := NewCapsuleParser(&buf) + ct, r, err := p.Next() + require.NoError(t, err) + require.Equal(t, CapsuleType(1337), ct) + val, err := io.ReadAll(r) + require.NoError(t, err) + require.Empty(t, val) + + ct, r, err = p.Next() + require.NoError(t, err) + require.Equal(t, CapsuleType(1337), ct) + val, err = io.ReadAll(r) + require.NoError(t, err) + require.Empty(t, val) +} diff --git a/third_party/quic-go/http3/client.go b/third_party/quic-go/http3/client.go new file mode 100644 index 0000000..021ca49 --- /dev/null +++ b/third_party/quic-go/http3/client.go @@ -0,0 +1,496 @@ +package http3 + +import ( + "context" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptrace" + "net/textproto" + "sync" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/quic-go/qpack" +) + +const ( + // MethodGet0RTT allows a GET request to be sent using 0-RTT. + // Note that 0-RTT doesn't provide replay protection and should only be used for idempotent requests. + MethodGet0RTT = "GET_0RTT" + // MethodHead0RTT allows a HEAD request to be sent using 0-RTT. + // Note that 0-RTT doesn't provide replay protection and should only be used for idempotent requests. + MethodHead0RTT = "HEAD_0RTT" +) + +const ( + defaultUserAgent = "quic-go HTTP/3" + defaultMaxResponseHeaderBytes = 10 * 1 << 20 // 10 MB +) + +var errGoAway = errors.New("connection in graceful shutdown") + +type errConnUnusable struct{ e error } + +func (e *errConnUnusable) Unwrap() error { return e.e } +func (e *errConnUnusable) Error() string { return fmt.Sprintf("http3: conn unusable: %s", e.e.Error()) } + +const max1xxResponses = 5 // arbitrary bound on number of informational responses + +var defaultQuicConfig = &quic.Config{ + MaxIncomingStreams: -1, // don't allow the server to create bidirectional streams + KeepAlivePeriod: 10 * time.Second, +} + +// ClientConn is an HTTP/3 client doing requests to a single remote server. +type ClientConn struct { + conn *quic.Conn + rawConn *rawConn + + decoder *qpack.Decoder + + // Additional HTTP/3 settings. + // It is invalid to specify any settings defined by RFC 9114 (HTTP/3) and RFC 9297 (HTTP Datagrams). + additionalSettings map[uint64]uint64 + + // maxResponseHeaderBytes specifies a limit on how many response bytes are + // allowed in the server's response header. + maxResponseHeaderBytes int + + // disableCompression, if true, prevents the Transport from requesting compression with an + // "Accept-Encoding: gzip" request header when the Request contains no existing Accept-Encoding value. + // If the Transport requests gzip on its own and gets a gzipped response, it's transparently + // decoded in the Response.Body. + // However, if the user explicitly requested gzip it is not automatically uncompressed. + disableCompression bool + + streamMx sync.Mutex + maxStreamID quic.StreamID // set once a GOAWAY frame is received + goAwayCtx context.Context + goAwayCancel context.CancelFunc + + qlogger qlogwriter.Recorder + logger *slog.Logger + + requestWriter *requestWriter +} + +var _ http.RoundTripper = &ClientConn{} + +func newClientConn( + conn *quic.Conn, + enableDatagrams bool, + additionalSettings map[uint64]uint64, + maxResponseHeaderBytes int, + disableCompression bool, + logger *slog.Logger, +) *ClientConn { + var qlogger qlogwriter.Recorder + if qlogTrace := conn.QlogTrace(); qlogTrace != nil && qlogTrace.SupportsSchemas(qlog.EventSchema) { + qlogger = qlogTrace.AddProducer() + } + c := &ClientConn{ + conn: conn, + additionalSettings: additionalSettings, + disableCompression: disableCompression, + maxStreamID: invalidStreamID, + logger: logger, + qlogger: qlogger, + decoder: qpack.NewDecoder(), + } + c.goAwayCtx, c.goAwayCancel = context.WithCancel(context.Background()) + if maxResponseHeaderBytes <= 0 { + c.maxResponseHeaderBytes = defaultMaxResponseHeaderBytes + } else { + c.maxResponseHeaderBytes = maxResponseHeaderBytes + } + c.requestWriter = newRequestWriter() + c.rawConn = newRawConn( + conn, + enableDatagrams, + c.onStreamsEmpty, + c.handleControlStream, + qlogger, + c.logger, + ) + // send the SETTINGs frame, using 0-RTT data, if possible + go func() { + _, err := c.rawConn.openControlStream(&settingsFrame{ + Datagram: enableDatagrams, + Other: additionalSettings, + MaxFieldSectionSize: int64(c.maxResponseHeaderBytes), + }) + if err != nil { + if c.logger != nil { + c.logger.Debug("setting up connection failed", "error", err) + } + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeInternalError), "") + return + } + }() + return c +} + +// OpenRequestStream opens a new request stream on the HTTP/3 connection. +func (c *ClientConn) OpenRequestStream(ctx context.Context) (*RequestStream, error) { + return c.openRequestStream(ctx, c.requestWriter, nil, c.disableCompression, c.maxResponseHeaderBytes) +} + +func (c *ClientConn) openRequestStream( + ctx context.Context, + requestWriter *requestWriter, + reqDone chan<- struct{}, + disableCompression bool, + maxHeaderBytes int, +) (*RequestStream, error) { + // RFC 9114 Section 5.2 prohibits opening any new request streams after GOAWAY. + // The stream ID only identifies requests that were already in flight and might still be processed. + if c.goAwayCtx.Err() != nil { + return nil, errGoAway + } + + openCtx, cancel := context.WithCancelCause(ctx) + // A request blocked in OpenStreamSync has no request stream yet, so it is not in flight. + stop := context.AfterFunc(c.goAwayCtx, func() { cancel(errGoAway) }) + str, err := c.conn.OpenStreamSync(openCtx) + stop() + cancel(nil) + if err != nil { + if context.Cause(openCtx) == errGoAway { + return nil, errGoAway + } + return nil, err + } + + // Check again in case GOAWAY raced with OpenStreamSync. + if c.goAwayCtx.Err() != nil { + str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled)) + return nil, errGoAway + } + + hstr := c.rawConn.TrackStream(str) + rsp := &http.Response{} + trace := httptrace.ContextClientTrace(ctx) + return newRequestStream( + newStream(hstr, c.rawConn, trace, func(r io.Reader, hf *headersFrame) error { + hdr, err := decodeTrailers(r, hf, maxHeaderBytes, c.decoder, c.qlogger, str.StreamID()) + if err != nil { + return err + } + rsp.Trailer = hdr + return nil + }, c.qlogger), + requestWriter, + reqDone, + c.decoder, + disableCompression, + maxHeaderBytes, + rsp, + ), nil +} + +func (c *ClientConn) handleUnidirectionalStream(str *quic.ReceiveStream) { + c.rawConn.handleUnidirectionalStream(str, false) +} + +func (c *ClientConn) handleControlStream(str *quic.ReceiveStream, fp *frameParser) { + for { + f, err := fp.ParseNext(c.qlogger) + if err != nil { + var serr *quic.StreamError + if err == io.EOF || errors.As(err, &serr) { + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeClosedCriticalStream), "") + return + } + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameError), "") + return + } + // GOAWAY is the only frame allowed at this point: + // * unexpected frames are ignored by the frame parser + // * we don't support any extension that might add support for more frames + goaway, ok := f.(*goAwayFrame) + if !ok { + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "") + return + } + if goaway.StreamID%4 != 0 { // client-initiated, bidirectional streams + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeIDError), "") + return + } + c.streamMx.Lock() + // the server is not allowed to increase the Stream ID in subsequent GOAWAY frames + if c.maxStreamID != invalidStreamID && goaway.StreamID > c.maxStreamID { + c.streamMx.Unlock() + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeIDError), "") + return + } + c.maxStreamID = goaway.StreamID + c.goAwayCancel() + c.streamMx.Unlock() + + hasActiveStreams := c.rawConn.hasActiveStreams() + // immediately close the connection if there are currently no active requests + if !hasActiveStreams { + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "") + return + } + } +} + +func (c *ClientConn) onStreamsEmpty() { + c.streamMx.Lock() + defer c.streamMx.Unlock() + + // The server is performing a graceful shutdown. + if c.maxStreamID != invalidStreamID { + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "") + } +} + +// RoundTrip executes a request and returns a response +func (c *ClientConn) RoundTrip(req *http.Request) (*http.Response, error) { + rsp, err := c.roundTrip(req) + if err != nil && req.Context().Err() != nil { + // if the context was canceled, return the context cancellation error + err = req.Context().Err() + } + return rsp, err +} + +func (c *ClientConn) roundTrip(req *http.Request) (*http.Response, error) { + // Immediately send out this request, if this is a 0-RTT request. + switch req.Method { + case MethodGet0RTT: + // don't modify the original request + reqCopy := *req + req = &reqCopy + req.Method = http.MethodGet + case MethodHead0RTT: + // don't modify the original request + reqCopy := *req + req = &reqCopy + req.Method = http.MethodHead + default: + // wait for the handshake to complete + select { + case <-c.conn.HandshakeComplete(): + case <-req.Context().Done(): + return nil, req.Context().Err() + } + } + + // It is only possible to send an Extended CONNECT request once the SETTINGS were received. + // See section 3 of RFC 8441. + if isExtendedConnectRequest(req) { + connCtx := c.conn.Context() + // wait for the server's SETTINGS frame to arrive + select { + case <-c.rawConn.ReceivedSettings(): + case <-connCtx.Done(): + return nil, context.Cause(connCtx) + } + if !c.rawConn.Settings().EnableExtendedConnect { + return nil, errors.New("http3: server didn't enable Extended CONNECT") + } + } + + reqDone := make(chan struct{}) + str, err := c.openRequestStream( + req.Context(), + c.requestWriter, + reqDone, + c.disableCompression, + c.maxResponseHeaderBytes, + ) + if err != nil { + return nil, &errConnUnusable{e: err} + } + + // Request Cancellation: + // This go routine keeps running even after RoundTripOpt() returns. + // It is shut down when the application is done processing the body. + done := make(chan struct{}) + go func() { + defer close(done) + select { + case <-req.Context().Done(): + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled)) + str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled)) + case <-reqDone: + } + }() + + rsp, err := c.doRequest(req, str) + if err != nil { // if any error occurred + close(reqDone) + <-done + return nil, maybeReplaceError(err) + } + return rsp, maybeReplaceError(err) +} + +// ReceivedSettings returns a channel that is closed once the server's HTTP/3 settings were received. +// Settings can be obtained from the Settings method after the channel was closed. +func (c *ClientConn) ReceivedSettings() <-chan struct{} { + return c.rawConn.ReceivedSettings() +} + +// Settings returns the HTTP/3 settings for this connection. +// It is only valid to call this function after the channel returned by ReceivedSettings was closed. +func (c *ClientConn) Settings() *Settings { + return c.rawConn.Settings() +} + +// CloseWithError closes the connection with the given error code and message. +// It is invalid to call this function after the connection was closed. +func (c *ClientConn) CloseWithError(code quic.ApplicationErrorCode, msg string) error { + return c.conn.CloseWithError(code, msg) +} + +// Context returns a context that is cancelled when the connection is closed. +func (c *ClientConn) Context() context.Context { + return c.conn.Context() +} + +// cancelingReader reads from the io.Reader. +// It cancels writing on the stream if any error other than io.EOF occurs. +type cancelingReader struct { + r io.Reader + str *RequestStream +} + +func (r *cancelingReader) Read(b []byte) (int, error) { + n, err := r.r.Read(b) + if err != nil && err != io.EOF { + r.str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled)) + } + return n, err +} + +func (c *ClientConn) sendRequestBody(str *RequestStream, body io.ReadCloser, contentLength int64) error { + defer body.Close() + buf := make([]byte, bodyCopyBufferSize) + sr := &cancelingReader{str: str, r: body} + if contentLength == -1 { + _, err := io.CopyBuffer(str, sr, buf) + return err + } + + // make sure we don't send more bytes than the content length + n, err := io.CopyBuffer(str, io.LimitReader(sr, contentLength), buf) + if err != nil { + return err + } + var extra int64 + extra, err = io.CopyBuffer(io.Discard, sr, buf) + n += extra + if n > contentLength { + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled)) + return fmt.Errorf("http: ContentLength=%d with Body length %d", contentLength, n) + } + return err +} + +func (c *ClientConn) doRequest(req *http.Request, str *RequestStream) (*http.Response, error) { + trace := httptrace.ContextClientTrace(req.Context()) + var sendingReqFailed bool + if err := str.sendRequestHeader(req); err != nil { + traceWroteRequest(trace, err) + if c.logger != nil { + c.logger.Debug("error writing request", "error", err) + } + sendingReqFailed = true + } + if !sendingReqFailed { + if req.Body == nil { + traceWroteRequest(trace, nil) + str.Close() + } else { + // send the request body asynchronously + go func() { + defer str.Close() + contentLength := int64(-1) + // According to the documentation for http.Request.ContentLength, + // a value of 0 with a non-nil Body is also treated as unknown content length. + if req.ContentLength > 0 { + contentLength = req.ContentLength + } + err := c.sendRequestBody(str, req.Body, contentLength) + traceWroteRequest(trace, err) + if err != nil { + if c.logger != nil { + c.logger.Debug("error writing request", "error", err) + } + return + } + + if len(req.Trailer) > 0 { + if err := str.sendRequestTrailer(req); err != nil { + if c.logger != nil { + c.logger.Debug("error writing trailers", "error", err) + } + } + } + }() + } + } + + // copy from net/http: support 1xx responses + var num1xx int // number of informational 1xx headers received + var res *http.Response + for { + var err error + res, err = str.ReadResponse() + if err != nil { + return nil, err + } + resCode := res.StatusCode + is1xx := 100 <= resCode && resCode <= 199 + // treat 101 as a terminal status, see https://github.com/golang/go/issues/26161 + is1xxNonTerminal := is1xx && resCode != http.StatusSwitchingProtocols + if is1xxNonTerminal { + num1xx++ + if num1xx > max1xxResponses { + str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + return nil, errors.New("http3: too many 1xx informational responses") + } + traceGot1xxResponse(trace, resCode, textproto.MIMEHeader(res.Header)) + if resCode == http.StatusContinue { + traceGot100Continue(trace) + } + continue + } + break + } + connState := c.conn.ConnectionState().TLS + res.TLS = &connState + res.Request = req + return res, nil +} + +// RawClientConn is a low-level HTTP/3 client connection. +// It allows the application to take control of the stream accept loops, +// giving the application the ability to handle streams originating from the server. +type RawClientConn struct { + *ClientConn +} + +// HandleUnidirectionalStream handles an incoming unidirectional stream. +func (c *RawClientConn) HandleUnidirectionalStream(str *quic.ReceiveStream) { + c.rawConn.handleUnidirectionalStream(str, false) +} + +// HandleBidirectionalStream handles an incoming bidirectional stream. +func (c *ClientConn) HandleBidirectionalStream(str *quic.Stream) { + // According to RFC 9114, the server is not allowed to open bidirectional streams. + c.rawConn.CloseWithError( + quic.ApplicationErrorCode(ErrCodeStreamCreationError), + fmt.Sprintf("server opened bidirectional stream %d", str.StreamID()), + ) +} diff --git a/third_party/quic-go/http3/client_test.go b/third_party/quic-go/http3/client_test.go new file mode 100644 index 0000000..1b45548 --- /dev/null +++ b/third_party/quic-go/http3/client_test.go @@ -0,0 +1,887 @@ +package http3 + +import ( + "bytes" + "compress/gzip" + "context" + "io" + mrand "math/rand/v2" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + "github.com/apernet/quic-go/testutils/events" + "github.com/quic-go/qpack" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func TestClientSettings(t *testing.T) { + t.Run("enable datagrams", func(t *testing.T) { + testClientSettings(t, true, nil) + }) + t.Run("additional settings", func(t *testing.T) { + testClientSettings(t, false, map[uint64]uint64{13: 37}) + }) +} + +func testClientSettings(t *testing.T, enableDatagrams bool, other map[uint64]uint64) { + tr := &Transport{ + EnableDatagrams: enableDatagrams, + AdditionalSettings: other, + } + + var eventRecorder events.Recorder + clientConn, serverConn := newConnPair(t, withClientRecorder(&eventRecorder)) + tr.NewClientConn(clientConn) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptUniStream(ctx) + require.NoError(t, err) + + str.SetReadDeadline(time.Now().Add(time.Second)) + typ, err := quicvarint.Read(quicvarint.NewReader(str)) + require.NoError(t, err) + require.EqualValues(t, streamTypeControlStream, typ) + fp := (&frameParser{r: str}) + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f) + settingsFrame := f.(*settingsFrame) + require.Equal(t, settingsFrame.Datagram, enableDatagrams) + require.Equal(t, settingsFrame.Other, other) + + var datagramValue *bool + if enableDatagrams { + datagramValue = pointer(true) + } + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameCreated{ + StreamID: str.StreamID(), + Raw: qlog.RawInfo{Length: 10}, + Frame: qlog.Frame{ + Frame: qlog.SettingsFrame{ + MaxFieldSectionSize: defaultMaxResponseHeaderBytes, + Datagram: datagramValue, + Other: other, + }, + }, + }, + }, + filterQlogEventsForFrame(eventRecorder.Events(qlog.FrameCreated{}), qlog.SettingsFrame{}), + ) +} + +func encodeResponse(t *testing.T, status int) []byte { + t.Helper() + + mockCtrl := gomock.NewController(t) + buf := &bytes.Buffer{} + rstr := NewMockDatagramStream(mockCtrl) + rstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + rstr.EXPECT().Write(gomock.Any()).Do(buf.Write).AnyTimes() + rw := newResponseWriter(newStream(rstr, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), nil, false, nil) + rw.WriteHeader(status) + rw.Flush() + return buf.Bytes() +} + +func TestClientRequest(t *testing.T) { + t.Run("GET", func(t *testing.T) { + rsp := testClientRequest(t, false, http.MethodGet, encodeResponse(t, http.StatusTeapot)) + require.Equal(t, http.StatusTeapot, rsp.StatusCode) + require.Equal(t, "HTTP/3.0", rsp.Proto) + require.Equal(t, 3, rsp.ProtoMajor) + require.NotNil(t, rsp.Request) + }) + + t.Run("GET 0-RTT", func(t *testing.T) { + rsp := testClientRequest(t, true, http.MethodGet, encodeResponse(t, http.StatusOK)) + require.Equal(t, http.StatusOK, rsp.StatusCode) + }) + + t.Run("HEAD", func(t *testing.T) { + rsp := testClientRequest(t, false, http.MethodHead, encodeResponse(t, http.StatusTeapot)) + require.Equal(t, http.StatusTeapot, rsp.StatusCode) + }) + + t.Run("HEAD 0-RTT", func(t *testing.T) { + rsp := testClientRequest(t, true, http.MethodHead, encodeResponse(t, http.StatusOK)) + require.Equal(t, http.StatusOK, rsp.StatusCode) + }) +} + +func testClientRequest(t *testing.T, use0RTT bool, method string, rspBytes []byte) *http.Response { + clientConn, serverConn := newConnPair(t) + + reqMethod := method + if use0RTT { + switch method { + case http.MethodGet: + reqMethod = MethodGet0RTT + case http.MethodHead: + reqMethod = MethodHead0RTT + } + } + req, err := http.NewRequest(reqMethod, "http://quic-go.net", nil) + require.NoError(t, err) + + type result struct { + rsp *http.Response + err error + } + resultChan := make(chan result, 1) + go func() { + cc := (&Transport{}).NewClientConn(clientConn) + rsp, err := cc.RoundTrip(req) + resultChan <- result{rsp: rsp, err: err} + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + str.SetReadDeadline(time.Now().Add(time.Second)) + + hfs := decodeHeader(t, str) + require.Equal(t, []string{method}, hfs[":method"]) + + _, err = str.Write(rspBytes) + require.NoError(t, err) + + var res result + select { + case res = <-resultChan: + require.NoError(t, res.err) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // make sure the http.Request.Method value was not modified + if use0RTT { + switch reqMethod { + case MethodGet0RTT: + require.Equal(t, req.Method, MethodGet0RTT) + case MethodHead0RTT: + require.Equal(t, req.Method, MethodHead0RTT) + } + } + return res.rsp +} + +func randomString(length int) string { + const alphabet = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" + b := make([]byte, length) + for i := range b { + n := mrand.IntN(len(alphabet)) + b[i] = alphabet[n] + } + return string(b) +} + +func TestClientRequestError(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + req, err := http.NewRequest(http.MethodGet, "http://quic-go.net", nil) + require.NoError(t, err) + for range 1000 { + req.Header.Add(randomString(50), randomString(50)) + } + + type result struct { + rsp *http.Response + err error + } + resultChan := make(chan result, 1) + go func() { + cc := (&Transport{}).NewClientConn(clientConn) + rsp, err := cc.RoundTrip(req) + resultChan <- result{rsp: rsp, err: err} + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + + _, err = str.Write(encodeResponse(t, http.StatusTeapot)) + require.NoError(t, err) + + var res result + select { + case res = <-resultChan: + require.NoError(t, res.err) + require.Equal(t, http.StatusTeapot, res.rsp.StatusCode) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestClientResponseValidation(t *testing.T) { + t.Run("HEADERS frame too large", func(t *testing.T) { + require.ErrorContains(t, + testClientResponseValidation(t, + &Transport{MaxResponseHeaderBytes: 1337}, + (&headersFrame{Length: 1338}).Append(nil), + quic.StreamErrorCode(ErrCodeFrameError), + ), + "http3: HEADERS frame too large", + ) + }) + + t.Run("invalid headers", func(t *testing.T) { + headerBuf := &bytes.Buffer{} + enc := qpack.NewEncoder(headerBuf) + // not a valid response pseudo header + require.NoError(t, enc.WriteField(qpack.HeaderField{Name: ":method", Value: "GET"})) + require.NoError(t, enc.Close()) + b := (&headersFrame{Length: uint64(headerBuf.Len())}).Append(nil) + b = append(b, headerBuf.Bytes()...) + + require.ErrorContains(t, + testClientResponseValidation(t, &Transport{}, b, quic.StreamErrorCode(ErrCodeMessageError)), + "invalid response pseudo header", + ) + }) +} + +func testClientResponseValidation(t *testing.T, tr *Transport, rsp []byte, expectedReset quic.StreamErrorCode) error { + clientConn, serverConn := newConnPair(t) + + cc := tr.NewClientConn(clientConn) + errChan := make(chan error) + go func() { + _, err := cc.RoundTrip(httptest.NewRequest(http.MethodGet, "http://quic-go.net", nil)) + errChan <- err + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + _, err = str.Write(rsp) + require.NoError(t, err) + + select { + case err := <-errChan: + expectStreamWriteReset(t, str, expectedReset) + // The client closes the stream after sending the request, + // so we need to wait for the RESET_STREAM frame to be received. + time.Sleep(scaleDuration(10 * time.Millisecond)) + expectStreamReadReset(t, str, expectedReset) + return err + case <-time.After(time.Second): + t.Fatal("timeout") + } + panic("unreachable") +} + +func TestClientRequestLengthLimit(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + cc := (&Transport{}).NewClientConn(clientConn) + errChan := make(chan error) + body := bytes.NewBufferString("request body") + go func() { + req := httptest.NewRequest(http.MethodPost, "http://quic-go.net", body) + req.ContentLength = 8 + _, err := cc.RoundTrip(req) + errChan <- err + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + + _, err = io.ReadAll(str) + var strErr *quic.StreamError + require.ErrorAs(t, err, &strErr) + require.Equal(t, quic.StreamErrorCode(ErrCodeRequestCanceled), strErr.ErrorCode) + + _, err = str.Write(encodeResponse(t, http.StatusTeapot)) + require.NoError(t, err) + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestClientExtendedConnect(t *testing.T) { + t.Run("enabled", func(t *testing.T) { + testClientExtendedConnect(t, true) + }) + + t.Run("disabled", func(t *testing.T) { + testClientExtendedConnect(t, false) + }) +} + +func testClientExtendedConnect(t *testing.T, enabled bool) { + clientConn, serverConn := newConnPair(t) + + cc := (&Transport{}).NewClientConn(clientConn) + req, err := http.NewRequest(http.MethodConnect, "http://quic-go.net", nil) + require.NoError(t, err) + req.Proto = "connect" + + errChan := make(chan error) + go func() { + _, err := cc.RoundTrip(req) + errChan <- err + }() + + select { + case <-errChan: + t.Fatal("RoundTrip should have blocked until SETTINGS were received") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + // now send the SETTINGS + settingsStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + settingsStr.SetWriteDeadline(time.Now().Add(time.Second)) + settingsFrame := &settingsFrame{ExtendedConnect: enabled} + _, err = settingsStr.Write(settingsFrame.Append(quicvarint.Append(nil, streamTypeControlStream))) + require.NoError(t, err) + + select { + case <-cc.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for settings") + } + settings := cc.Settings() + require.Equal(t, enabled, settings.EnableExtendedConnect) + + if enabled { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + str.CancelRead(1337) + str.CancelWrite(1337) + } + + select { + case err := <-errChan: + if enabled { + require.ErrorIs(t, err, &Error{Remote: true, ErrorCode: 1337}) + } else { + require.EqualError(t, err, "http3: server didn't enable Extended CONNECT") + } + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestClient1xxHandling(t *testing.T) { + t.Run("a few early hints", func(t *testing.T) { + testClient1xxHandling(t, max1xxResponses, http.StatusOK, false) + }) + t.Run("too many early hints", func(t *testing.T) { + testClient1xxHandling(t, max1xxResponses+1, http.StatusOK, true) + }) + t.Run("EarlyHints followed by StatusSwitchingProtocols", func(t *testing.T) { + testClient1xxHandling(t, 1, http.StatusSwitchingProtocols, false) + }) +} + +func testClient1xxHandling(t *testing.T, numEarlyHints int, terminalStatus int, tooMany bool) { + var rspBuf bytes.Buffer + rstr := NewMockDatagramStream(gomock.NewController(t)) + rstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + rstr.EXPECT().Write(gomock.Any()).Do(rspBuf.Write).AnyTimes() + rw := newResponseWriter(newStream(rstr, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), nil, false, nil) + rw.header.Add("Link", "foo") + rw.header.Add("Link", "bar") + for range numEarlyHints { + rw.WriteHeader(http.StatusEarlyHints) + } + rw.WriteHeader(terminalStatus) + rw.Flush() + rspBytes := rspBuf.Bytes() + + clientConn, serverConn := newConnPair(t) + + type result struct { + rsp *http.Response + err error + } + resultChan := make(chan result, 1) + go func() { + cc := (&Transport{}).NewClientConn(clientConn) + rsp, err := cc.RoundTrip(httptest.NewRequest(http.MethodGet, "http://quic-go.net", nil)) + resultChan <- result{rsp: rsp, err: err} + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + + // request headers + hfs := decodeHeader(t, str) + require.Equal(t, hfs[":method"], []string{http.MethodGet}) + + _, err = str.Write(rspBytes) + require.NoError(t, err) + + var rsp *http.Response + select { + case res := <-resultChan: + if tooMany { + require.EqualError(t, res.err, "http3: too many 1xx informational responses") + return + } + require.NoError(t, res.err) + rsp = res.rsp + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Equal(t, []string{"foo", "bar"}, rsp.Header["Link"]) + require.Equal(t, terminalStatus, rsp.StatusCode) +} + +func TestClientGzip(t *testing.T) { + var buf bytes.Buffer + w := gzip.NewWriter(&buf) + w.Write([]byte("foobar")) + w.Close() + gzippedFoobar := buf.Bytes() + + t.Run("gzipped", func(t *testing.T) { + testClientGzip(t, gzippedFoobar, []byte("foobar"), false, true) + }) + t.Run("not gzipped", func(t *testing.T) { + testClientGzip(t, []byte("foobar"), []byte("foobar"), false, false) + }) + t.Run("disable compression", func(t *testing.T) { + testClientGzip(t, gzippedFoobar, gzippedFoobar, true, true) + }) +} + +func testClientGzip(t *testing.T, + data []byte, + expectedRsp []byte, + transportDisableCompression bool, + responseAddContentEncoding bool, +) { + var rspBuf bytes.Buffer + rstr := NewMockDatagramStream(gomock.NewController(t)) + rstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + rstr.EXPECT().Write(gomock.Any()).Do(rspBuf.Write).AnyTimes() + rw := newResponseWriter(newStream(rstr, nil, nil, func(io.Reader, *headersFrame) error { return nil }, nil), nil, false, nil) + rw.WriteHeader(http.StatusOK) + if responseAddContentEncoding { + rw.header.Add("Content-Encoding", "gzip") + } + rw.Write(data) + rw.Flush() + + clientConn, serverConn := newConnPair(t) + + type result struct { + rsp *http.Response + err error + } + resultChan := make(chan result) + go func() { + cc := (&Transport{DisableCompression: transportDisableCompression}).NewClientConn(clientConn) + rsp, err := cc.RoundTrip(httptest.NewRequest(http.MethodGet, "http://quic-go.net", nil)) + resultChan <- result{rsp: rsp, err: err} + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + + // request headers + str.SetReadDeadline(time.Now().Add(time.Second)) + hfs := decodeHeader(t, str) + if transportDisableCompression { + require.NotContains(t, hfs, "accept-encoding") + } else { + require.Equal(t, hfs["accept-encoding"], []string{"gzip"}) + } + + _, err = str.Write(rspBuf.Bytes()) + require.NoError(t, err) + require.NoError(t, str.Close()) + + var rsp *http.Response + select { + case res := <-resultChan: + require.NoError(t, res.err) + rsp = res.rsp + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Equal(t, http.StatusOK, rsp.StatusCode) + body, err := io.ReadAll(rsp.Body) + require.NoError(t, err) + require.Equal(t, expectedRsp, body) +} + +func TestClientRequestCancellation(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + requestCtx, requestCancel := context.WithCancel(context.Background()) + req, err := http.NewRequestWithContext(requestCtx, http.MethodGet, "http://quic-go.net", nil) + require.NoError(t, err) + + type result struct { + rsp *http.Response + err error + } + resultChan := make(chan result) + go func() { + cc := (&Transport{}).NewClientConn(clientConn) + rsp, err := cc.RoundTrip(req) + resultChan <- result{rsp: rsp, err: err} + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + + _, err = str.Write(encodeResponse(t, http.StatusTeapot)) + require.NoError(t, err) + + select { + case res := <-resultChan: + require.NoError(t, res.err) + require.Equal(t, http.StatusTeapot, res.rsp.StatusCode) + case <-time.After(time.Second): + t.Fatal("timeout") + } + requestCancel() + + expectStreamWriteReset(t, str, quic.StreamErrorCode(ErrCodeRequestCanceled)) +} + +func TestClientConnGoAway(t *testing.T) { + t.Run("no active streams", func(t *testing.T) { + testClientConnGoAway(t, false) + }) + + t.Run("active stream", func(t *testing.T) { + testClientConnGoAway(t, true) + }) +} + +func testClientConnGoAway(t *testing.T, withStream bool) { + var clientEventRecorder events.Recorder + clientConn, serverConn := newConnPair(t, withClientRecorder(&clientEventRecorder)) + + cc := (&Transport{}).NewClientConn(clientConn) + + var str *RequestStream + if withStream { + s, err := cc.OpenRequestStream(context.Background()) + require.NoError(t, err) + str = s + } + + // server sends control stream with SETTINGS and GOAWAY + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{}).Append(b) + b = (&goAwayFrame{StreamID: 8}).Append(b) + controlStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(b) + require.NoError(t, err) + + // the connection should be closed after the stream is closed + if withStream { + select { + case <-serverConn.Context().Done(): + t.Fatal("connection closed") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + // GOAWAY allows the request that was already open to finish. + str.Close() + str.CancelRead(1337) + } + + select { + case <-serverConn.Context().Done(): + require.ErrorIs(t, + context.Cause(serverConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeNoError)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } + + expectedLen, expectedPayloadLen := expectedFrameLength(t, &goAwayFrame{StreamID: 8}) + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameParsed{ + StreamID: controlStr.StreamID(), + Raw: qlog.RawInfo{PayloadLength: expectedPayloadLen, Length: expectedLen}, + Frame: qlog.Frame{Frame: qlog.GoAwayFrame{StreamID: 8}}, + }, + }, + filterQlogEventsForFrame(clientEventRecorder.Events(qlog.FrameParsed{}), qlog.GoAwayFrame{StreamID: 8}), + ) +} + +func TestClientConnGoAwayConcurrent(t *testing.T) { + clientConn, serverConn := newConnPair(t, withServerBidiStreamLimit(2)) // allows streams 0 and 4 + + cc := (&Transport{}).NewClientConn(clientConn) + + // peer sends control stream with SETTINGS, but not GOAWAY yet + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{}).Append(b) + controlStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(b) + require.NoError(t, err) + + select { + case <-serverConn.Context().Done(): + t.Fatal("connection closed") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + // Consume both streams the server allows, but keep their receive sides open. + for range 2 { + str, err := cc.OpenRequestStream(context.Background()) + require.NoError(t, err) + require.NoError(t, str.Close()) + } + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + sstr0, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + require.Equal(t, quic.StreamID(0), sstr0.StreamID()) + sstr4, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + require.Equal(t, quic.StreamID(4), sstr4.StreamID()) + + // Both of these calls will block in OpenStreamSync. + errChan := make(chan error, 2) + for range 2 { + go func() { + str, err := cc.OpenRequestStream(context.Background()) + if err == nil { + str.Close() + } + errChan <- err + }() + } + + select { + case <-errChan: + t.Fatal("OpenStreamSync calls should have blocked") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + // Send a GOAWAY with a stream ID higher than the next stream ID. + b = (&goAwayFrame{StreamID: 12}).Append(nil) + _, err = controlStr.Write(b) + require.NoError(t, err) + + // The GOAWAY stream ID only applies to requests already in flight. Even though + // stream 8 would be below the limit, no new request stream may be opened. + for range 2 { + select { + case err := <-errChan: + require.ErrorIs(t, err, errGoAway) + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + + // The streams opened before GOAWAY are not canceled. + buf := make([]byte, 1) + _, err = sstr0.Read(buf) + require.ErrorIs(t, err, io.EOF) + _, err = sstr4.Read(buf) + require.ErrorIs(t, err, io.EOF) + + // Complete stream 0 to free stream credit, while stream 4 keeps the connection alive. + require.NoError(t, sstr0.Close()) + + // Calls made after receiving GOAWAY fail even when stream credit is available. + openCtx, openCancel := context.WithTimeout(context.Background(), time.Second) + defer openCancel() + _, err = cc.OpenRequestStream(openCtx) + require.ErrorIs(t, err, errGoAway) + + // In particular, the client didn't open and cancel stream 8 behind the caller's back. + noStreamCtx, noStreamCancel := context.WithTimeout(context.Background(), scaleDuration(10*time.Millisecond)) + defer noStreamCancel() + _, err = serverConn.AcceptStream(noStreamCtx) + require.ErrorIs(t, err, context.DeadlineExceeded) + + require.NoError(t, sstr4.Close()) +} + +func TestClientConnGoAwayFailures(t *testing.T) { + t.Run("invalid frame", func(t *testing.T) { + b := (&settingsFrame{}).Append(nil) + // 1337 is invalid value for the Extended CONNECT setting + b = (&settingsFrame{Other: map[uint64]uint64{settingExtendedConnect: 1337}}).Append(b) + testClientConnGoAwayFailures(t, b, nil, ErrCodeFrameError) + }) + + t.Run("not a GOAWAY", func(t *testing.T) { + b := (&settingsFrame{}).Append(nil) + // GOAWAY is the only allowed frame type after SETTINGS + b = (&headersFrame{}).Append(b) + testClientConnGoAwayFailures(t, b, nil, ErrCodeFrameUnexpected) + }) + + t.Run("stream closed before GOAWAY", func(t *testing.T) { + testClientConnGoAwayFailures(t, (&settingsFrame{}).Append(nil), io.EOF, ErrCodeClosedCriticalStream) + }) + + t.Run("stream reset before GOAWAY", func(t *testing.T) { + testClientConnGoAwayFailures(t, + (&settingsFrame{}).Append(nil), + &quic.StreamError{Remote: true, ErrorCode: 42}, + ErrCodeClosedCriticalStream, + ) + }) + + t.Run("invalid stream ID", func(t *testing.T) { + data := (&settingsFrame{}).Append(nil) + data = (&goAwayFrame{StreamID: 1}).Append(data) + testClientConnGoAwayFailures(t, data, nil, ErrCodeIDError) + }) + + t.Run("increased stream ID", func(t *testing.T) { + localConn, peerConn := newConnPair(t) + + cc := (&Transport{}).NewClientConn(localConn) + + // need an active stream so the connection doesn't close after the first GOAWAY + _, err := cc.OpenRequestStream(context.Background()) + require.NoError(t, err) + + controlStr, err := peerConn.OpenUniStream() + require.NoError(t, err) + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{}).Append(b) + b = (&goAwayFrame{StreamID: 4}).Append(b) + b = (&goAwayFrame{StreamID: 8}).Append(b) + _, err = controlStr.Write(b) + require.NoError(t, err) + + select { + case <-peerConn.Context().Done(): + require.ErrorIs(t, + context.Cause(peerConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeIDError)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } + }) +} + +func testClientConnGoAwayFailures(t *testing.T, data []byte, readErr error, expectedErr ErrCode) { + localConn, peerConn := newConnPair(t) + + (&Transport{}).NewClientConn(localConn) + + controlStr, err := peerConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(quicvarint.Append(nil, streamTypeControlStream)) + require.NoError(t, err) + + switch readErr { + case nil: + _, err = controlStr.Write(data) + require.NoError(t, err) + case io.EOF: + _, err = controlStr.Write(data) + require.NoError(t, err) + require.NoError(t, controlStr.Close()) + default: + // make sure the stream type is received + time.Sleep(scaleDuration(10 * time.Millisecond)) + controlStr.CancelWrite(1337) + } + + select { + case <-peerConn.Context().Done(): + require.ErrorIs(t, + context.Cause(peerConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(expectedErr)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } +} + +func TestClientConnHandleBidirectionalStream(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + cc := (&Transport{}).NewClientConn(clientConn) + + str, err := clientConn.OpenStream() + require.NoError(t, err) + cc.HandleBidirectionalStream(str) + + select { + case <-serverConn.Context().Done(): + require.ErrorIs(t, + context.Cause(serverConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeStreamCreationError)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for connection close") + } +} + +func TestRawClientConnHandleUnidirectionalStream(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + cc := (&Transport{}).NewRawClientConn(clientConn) + + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{}).Append(b) + str, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + uniStr, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + cc.HandleUnidirectionalStream(uniStr) + }() + + select { + case <-cc.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for settings") + } + require.NotNil(t, cc.Settings()) +} diff --git a/third_party/quic-go/http3/conn.go b/third_party/quic-go/http3/conn.go new file mode 100644 index 0000000..c50ae3a --- /dev/null +++ b/third_party/quic-go/http3/conn.go @@ -0,0 +1,319 @@ +package http3 + +import ( + "context" + "errors" + "fmt" + "io" + "log/slog" + "maps" + "net" + "sync" + "sync/atomic" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" +) + +const maxQuarterStreamID = 1<<60 - 1 + +// invalidStreamID is a stream ID that is invalid. The first valid stream ID in QUIC is 0. +const invalidStreamID = quic.StreamID(-1) + +// rawConn is an HTTP/3 connection. +// It provides HTTP/3 specific functionality by wrapping a quic.Conn, +// in particular handling of unidirectional HTTP/3 streams, SETTINGS and datagrams. +type rawConn struct { + conn *quic.Conn + + logger *slog.Logger + + enableDatagrams bool + + streamMx sync.Mutex + streams map[quic.StreamID]*stateTrackingStream + + rcvdControlStr atomic.Bool + rcvdQPACKEncoderStr atomic.Bool + rcvdQPACKDecoderStr atomic.Bool + controlStrHandler func(*quic.ReceiveStream, *frameParser) // is called *after* the SETTINGS frame was parsed + + onStreamsEmpty func() + + settings *Settings + receivedSettings chan struct{} + + qlogger qlogwriter.Recorder + qloggerWG sync.WaitGroup // tracks goroutines that may produce qlog events +} + +func newRawConn( + quicConn *quic.Conn, + enableDatagrams bool, + onStreamsEmpty func(), + controlStrHandler func(*quic.ReceiveStream, *frameParser), + qlogger qlogwriter.Recorder, + logger *slog.Logger, +) *rawConn { + c := &rawConn{ + conn: quicConn, + logger: logger, + enableDatagrams: enableDatagrams, + receivedSettings: make(chan struct{}), + streams: make(map[quic.StreamID]*stateTrackingStream), + qlogger: qlogger, + onStreamsEmpty: onStreamsEmpty, + controlStrHandler: controlStrHandler, + } + if qlogger != nil { + context.AfterFunc(quicConn.Context(), c.closeQlogger) + } + return c +} + +func (c *rawConn) OpenUniStream() (*quic.SendStream, error) { + return c.conn.OpenUniStream() +} + +// openControlStream opens the control stream and sends the SETTINGS frame. +// It returns the control stream (needed by the server for sending GOAWAY later). +func (c *rawConn) openControlStream(settings *settingsFrame) (*quic.SendStream, error) { + c.qloggerWG.Add(1) + defer c.qloggerWG.Done() + + str, err := c.conn.OpenUniStream() + if err != nil { + return nil, err + } + b := make([]byte, 0, 64) + b = quicvarint.Append(b, streamTypeControlStream) + b = settings.Append(b) + if c.qlogger != nil { + sf := qlog.SettingsFrame{ + MaxFieldSectionSize: settings.MaxFieldSectionSize, + Other: maps.Clone(settings.Other), + } + if settings.Datagram { + sf.Datagram = pointer(true) + } + if settings.ExtendedConnect { + sf.ExtendedConnect = pointer(true) + } + c.qlogger.RecordEvent(qlog.FrameCreated{ + StreamID: str.StreamID(), + Raw: qlog.RawInfo{Length: len(b)}, + Frame: qlog.Frame{Frame: sf}, + }) + } + if _, err := str.Write(b); err != nil { + return nil, err + } + return str, nil +} + +func (c *rawConn) TrackStream(str *quic.Stream) *stateTrackingStream { + hstr := newStateTrackingStream(str, c, func(b []byte) error { return c.sendDatagram(str.StreamID(), b) }) + + c.streamMx.Lock() + c.streams[str.StreamID()] = hstr + c.qloggerWG.Add(1) + c.streamMx.Unlock() + return hstr +} + +func (c *rawConn) RemoteAddr() net.Addr { + return c.conn.RemoteAddr() +} + +func (c *rawConn) ConnectionState() quic.ConnectionState { + return c.conn.ConnectionState() +} + +func (c *rawConn) clearStream(id quic.StreamID) { + c.streamMx.Lock() + defer c.streamMx.Unlock() + + if _, ok := c.streams[id]; ok { + delete(c.streams, id) + c.qloggerWG.Done() + } + if len(c.streams) == 0 { + c.onStreamsEmpty() + } +} + +func (c *rawConn) hasActiveStreams() bool { + c.streamMx.Lock() + defer c.streamMx.Unlock() + + return len(c.streams) > 0 +} + +func (c *rawConn) CloseWithError(code quic.ApplicationErrorCode, msg string) error { + return c.conn.CloseWithError(code, msg) +} + +func (c *rawConn) handleUnidirectionalStream(str *quic.ReceiveStream, isServer bool) { + c.qloggerWG.Add(1) + defer c.qloggerWG.Done() + + streamType, err := quicvarint.Read(quicvarint.NewReader(str)) + if err != nil { + if c.logger != nil { + c.logger.Debug("reading stream type on stream failed", "stream ID", str.StreamID(), "error", err) + } + return + } + // We're only interested in the control stream here. + switch streamType { + case streamTypeControlStream: + case streamTypeQPACKEncoderStream: + if isFirst := c.rcvdQPACKEncoderStr.CompareAndSwap(false, true); !isFirst { + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "duplicate QPACK encoder stream") + } + // Our QPACK implementation doesn't use the dynamic table yet. + return + case streamTypeQPACKDecoderStream: + if isFirst := c.rcvdQPACKDecoderStr.CompareAndSwap(false, true); !isFirst { + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "duplicate QPACK decoder stream") + } + // Our QPACK implementation doesn't use the dynamic table yet. + return + case streamTypePushStream: + if isServer { + // only the server can push + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "") + } else { + // we never increased the Push ID, so we don't expect any push streams + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeIDError), "") + } + return + default: + str.CancelRead(quic.StreamErrorCode(ErrCodeStreamCreationError)) + return + } + // Only a single control stream is allowed. + if isFirstControlStr := c.rcvdControlStr.CompareAndSwap(false, true); !isFirstControlStr { + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "duplicate control stream") + return + } + c.handleControlStream(str) +} + +func (c *rawConn) handleControlStream(str *quic.ReceiveStream) { + fp := &frameParser{closeConn: c.conn.CloseWithError, r: str, streamID: str.StreamID()} + f, err := fp.ParseNext(c.qlogger) + if err != nil { + var serr *quic.StreamError + if err == io.EOF || errors.As(err, &serr) { + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeClosedCriticalStream), "") + return + } + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameError), "") + return + } + sf, ok := f.(*settingsFrame) + if !ok { + c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeMissingSettings), "") + return + } + c.settings = &Settings{ + EnableDatagrams: sf.Datagram, + EnableExtendedConnect: sf.ExtendedConnect, + Other: sf.Other, + } + close(c.receivedSettings) + if sf.Datagram { + // If datagram support was enabled on our side as well as on the server side, + // we can expect it to have been negotiated both on the transport and on the HTTP/3 layer. + // Note: ConnectionState() will block until the handshake is complete (relevant when using 0-RTT). + if c.enableDatagrams && !c.ConnectionState().SupportsDatagrams.Remote { + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeSettingsError), "missing QUIC Datagram support") + return + } + c.qloggerWG.Go(func() { + if err := c.receiveDatagrams(); err != nil { + if c.logger != nil { + c.logger.Debug("receiving datagrams failed", "error", err) + } + } + }) + } + + if c.controlStrHandler != nil { + c.controlStrHandler(str, fp) + } +} + +func (c *rawConn) sendDatagram(streamID quic.StreamID, b []byte) error { + // TODO: this creates a lot of garbage and an additional copy + data := make([]byte, 0, len(b)+8) + quarterStreamID := uint64(streamID / 4) + data = quicvarint.Append(data, uint64(streamID/4)) + data = append(data, b...) + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.DatagramCreated{ + QuarterStreamID: quarterStreamID, + Raw: qlog.RawInfo{ + Length: len(data), + PayloadLength: len(b), + }, + }) + } + return c.conn.SendDatagram(data) +} + +func (c *rawConn) receiveDatagrams() error { + for { + b, err := c.conn.ReceiveDatagram(context.Background()) + if err != nil { + return err + } + quarterStreamID, n, err := quicvarint.Parse(b) + if err != nil { + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeDatagramError), "") + return fmt.Errorf("could not read quarter stream id: %w", err) + } + if c.qlogger != nil { + c.qlogger.RecordEvent(qlog.DatagramParsed{ + QuarterStreamID: quarterStreamID, + Raw: qlog.RawInfo{ + Length: len(b), + PayloadLength: len(b) - n, + }, + }) + } + if quarterStreamID > maxQuarterStreamID { + c.CloseWithError(quic.ApplicationErrorCode(ErrCodeDatagramError), "") + return fmt.Errorf("invalid quarter stream id: %w", err) + } + streamID := quic.StreamID(4 * quarterStreamID) + c.streamMx.Lock() + dg, ok := c.streams[streamID] + c.streamMx.Unlock() + if !ok { + continue + } + dg.enqueueDatagram(b[n:]) + } +} + +// ReceivedSettings returns a channel that is closed once the peer's SETTINGS frame was received. +// Settings can be optained from the Settings method after the channel was closed. +func (c *rawConn) ReceivedSettings() <-chan struct{} { return c.receivedSettings } + +// Settings returns the settings received on this connection. +// It is only valid to call this function after the channel returned by ReceivedSettings was closed. +func (c *rawConn) Settings() *Settings { return c.settings } + +// closeQlogger waits for all goroutines that may produce qlog events to finish, +// then closes the qlogger. +func (c *rawConn) closeQlogger() { + if c.qlogger == nil { + return + } + c.qloggerWG.Wait() + c.qlogger.Close() +} diff --git a/third_party/quic-go/http3/conn_test.go b/third_party/quic-go/http3/conn_test.go new file mode 100644 index 0000000..eed8779 --- /dev/null +++ b/third_party/quic-go/http3/conn_test.go @@ -0,0 +1,502 @@ +package http3 + +import ( + "bytes" + "context" + "io" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func TestConnReceiveSettings(t *testing.T) { + var eventRecorder events.Recorder + clientConn, serverConn := newConnPair(t, withServerRecorder(&eventRecorder)) + + conn := newRawConn(serverConn, false, nil, nil, &eventRecorder, nil) + b := quicvarint.Append(nil, streamTypeControlStream) + sf := &settingsFrame{ + MaxFieldSectionSize: 1234, + Datagram: true, + ExtendedConnect: true, + Other: map[uint64]uint64{1337: 42}, + } + b = sf.Append(b) + controlStr, err := clientConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverStr, err := serverConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(serverStr, true) + }() + select { + case <-conn.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for settings") + } + settings := conn.Settings() + require.True(t, settings.EnableDatagrams) + require.True(t, settings.EnableExtendedConnect) + require.Equal(t, map[uint64]uint64{1337: 42}, settings.Other) + + expectedLen, expectedPayloadLen := expectedFrameLength(t, sf) + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameParsed{ + StreamID: controlStr.StreamID(), + Raw: qlog.RawInfo{Length: expectedLen, PayloadLength: expectedPayloadLen}, + Frame: qlog.Frame{ + Frame: qlog.SettingsFrame{ + MaxFieldSectionSize: 1234, + Datagram: pointer(true), + ExtendedConnect: pointer(true), + Other: map[uint64]uint64{1337: 42}, + }, + }, + }, + }, + filterQlogEventsForFrame(eventRecorder.Events(qlog.FrameParsed{}), qlog.SettingsFrame{}), + ) +} + +func TestConnRejectDuplicateStreams(t *testing.T) { + t.Run("control stream", func(t *testing.T) { + testConnRejectDuplicateStreams(t, streamTypeControlStream) + }) + t.Run("encoder stream", func(t *testing.T) { + testConnRejectDuplicateStreams(t, streamTypeQPACKEncoderStream) + }) + t.Run("decoder stream", func(t *testing.T) { + testConnRejectDuplicateStreams(t, streamTypeQPACKDecoderStream) + }) +} + +func testConnRejectDuplicateStreams(t *testing.T, typ uint64) { + clientConn, serverConn := newConnPair(t) + + conn := newRawConn(serverConn, false, nil, nil, nil, nil) + b := quicvarint.Append(nil, typ) + if typ == streamTypeControlStream { + b = (&settingsFrame{}).Append(b) + } + controlStr1, err := clientConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr1.Write(b) + require.NoError(t, err) + controlStr2, err := clientConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr2.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverStr1, err := serverConn.AcceptUniStream(ctx) + require.NoError(t, err) + serverStr2, err := serverConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}, 2) + go func() { + defer func() { done <- struct{}{} }() + conn.handleUnidirectionalStream(serverStr1, true) + }() + go func() { + defer func() { done <- struct{}{} }() + conn.handleUnidirectionalStream(serverStr2, true) + }() + select { + case <-clientConn.Context().Done(): + require.ErrorIs(t, + context.Cause(clientConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeStreamCreationError)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for duplicate stream") + } + for range 2 { + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + } +} + +func TestConnResetUnknownUniStream(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + conn := newRawConn(serverConn, false, nil, nil, nil, nil) + buf := bytes.NewBuffer(quicvarint.Append(nil, 0x1337)) + str, err := clientConn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write(buf.Bytes()) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverStr, err := serverConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(serverStr, true) + }() + expectStreamWriteReset(t, str, quic.StreamErrorCode(ErrCodeStreamCreationError)) + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnControlStreamFailures(t *testing.T) { + t.Run("missing SETTINGS", func(t *testing.T) { + testConnControlStreamFailures(t, (&dataFrame{}).Append(nil), nil, ErrCodeMissingSettings) + }) + t.Run("frame error", func(t *testing.T) { + testConnControlStreamFailures(t, + // 1337 is invalid value for the Extended CONNECT setting + (&settingsFrame{Other: map[uint64]uint64{settingExtendedConnect: 1337}}).Append(nil), + nil, + ErrCodeFrameError, + ) + }) + t.Run("control stream closed before SETTINGS", func(t *testing.T) { + testConnControlStreamFailures(t, nil, io.EOF, ErrCodeClosedCriticalStream) + }) + t.Run("control stream reset before SETTINGS", func(t *testing.T) { + testConnControlStreamFailures(t, + nil, + &quic.StreamError{Remote: true, ErrorCode: 42}, + ErrCodeClosedCriticalStream, + ) + }) +} + +func testConnControlStreamFailures(t *testing.T, data []byte, readErr error, expectedErr ErrCode) { + clientConn, serverConn := newConnPair(t) + + conn := newRawConn(clientConn, false, nil, nil, nil, nil) + controlStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(quicvarint.Append(nil, streamTypeControlStream)) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientStr, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(clientStr, false) + }() + + switch readErr { + case nil: + _, err = controlStr.Write(data) + require.NoError(t, err) + case io.EOF: + _, err = controlStr.Write(data) + require.NoError(t, err) + require.NoError(t, controlStr.Close()) + default: + // make sure the stream type is received + time.Sleep(scaleDuration(10 * time.Millisecond)) + controlStr.CancelWrite(1337) + } + + select { + case <-serverConn.Context().Done(): + require.ErrorIs(t, + context.Cause(serverConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(expectedErr)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnControlStreamHandler(t *testing.T) { + t.Run("with handler", func(t *testing.T) { testConnControlStreamHandler(t, true) }) + t.Run("without handler", func(t *testing.T) { testConnControlStreamHandler(t, false) }) +} + +func testConnControlStreamHandler(t *testing.T, useHandler bool) { + localConn, peerConn := newConnPair(t) + + handlerCalled := make(chan struct{}) + var controlStrHandler func(*quic.ReceiveStream, *frameParser) + if useHandler { + controlStrHandler = func(*quic.ReceiveStream, *frameParser) { close(handlerCalled) } + } + conn := newRawConn(localConn, false, nil, controlStrHandler, nil, nil) + + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{}).Append(b) + str, err := peerConn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + localStr, err := localConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(localStr, false) + }() + + select { + case <-conn.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for settings") + } + if useHandler { + select { + case <-handlerCalled: + case <-time.After(time.Second): + t.Fatal("timeout waiting for handler to be called") + } + } else { + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout waiting for handler to return") + } + } +} + +func TestConnRejectPushStream(t *testing.T) { + t.Run("client", func(t *testing.T) { + testConnRejectPushStream(t, false, ErrCodeIDError) + }) + t.Run("server", func(t *testing.T) { + testConnRejectPushStream(t, true, ErrCodeStreamCreationError) + }) +} + +func testConnRejectPushStream(t *testing.T, isServer bool, expectedErr ErrCode) { + localConn, peerConn := newConnPair(t) + + conn := newRawConn(localConn, false, nil, nil, nil, nil) + buf := bytes.NewBuffer(quicvarint.Append(nil, streamTypePushStream)) + str, err := peerConn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write(buf.Bytes()) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + localStr, err := localConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(localStr, isServer) + }() + select { + case <-peerConn.Context().Done(): + require.ErrorIs(t, + context.Cause(peerConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(expectedErr)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestConnInconsistentDatagramSupport(t *testing.T) { + clientConn, serverConn := newConnPair(t) + + conn := newRawConn(clientConn, true, nil, nil, nil, nil) + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{Datagram: true}).Append(b) + controlStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientStr, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(clientStr, false) + }() + + select { + case <-serverConn.Context().Done(): + err := context.Cause(serverConn.Context()) + require.ErrorIs(t, err, &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeSettingsError)}) + require.ErrorContains(t, err, "missing QUIC Datagram support") + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } +} + +func TestConnSendAndReceiveDatagram(t *testing.T) { + var eventRecorder events.Recorder + clientConn, serverConn := newConnPair(t, withDatagrams(), withClientRecorder(&eventRecorder)) + + conn := newRawConn(clientConn, true, nil, nil, &eventRecorder, nil) + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{Datagram: true}).Append(b) + controlStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientStr, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + conn.handleUnidirectionalStream(clientStr, false) + }() + + const strID = 4 + + // first deliver a datagram... + // since the stream is not open yet, it will be dropped + quarterStreamID := quicvarint.Append([]byte{}, strID/4) + + datagram := append(quarterStreamID, []byte("foo")...) + require.NoError(t, serverConn.SendDatagram(datagram)) + time.Sleep(scaleDuration(10 * time.Millisecond)) // give the datagram a chance to be delivered + + require.Equal(t, + []qlogwriter.Event{ + qlog.DatagramParsed{ + QuarterStreamID: strID / 4, + Raw: qlog.RawInfo{Length: len(datagram), PayloadLength: 3}, + }, + }, + eventRecorder.Events(qlog.DatagramParsed{}), + ) + eventRecorder.Clear() + + // don't use stream 0, since that makes it hard to test that the quarter stream ID is used + str0, err := clientConn.OpenStreamSync(context.Background()) + require.NoError(t, err) + str0.Close() + + str, err := clientConn.OpenStream() + require.NoError(t, err) + require.Equal(t, quic.StreamID(strID), str.StreamID()) + datagramStr := conn.TrackStream(str) + + // now open the stream... + require.NoError(t, serverConn.SendDatagram(append(quarterStreamID, []byte("bar")...))) + + data, err := datagramStr.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("bar"), data) + + // now send a datagram + require.NoError(t, datagramStr.SendDatagram([]byte("foobaz"))) + + expected := quicvarint.Append([]byte{}, strID/4) + expected = append(expected, []byte("foobaz")...) + + require.Equal(t, + []qlogwriter.Event{ + qlog.DatagramCreated{ + QuarterStreamID: strID / 4, + Raw: qlog.RawInfo{PayloadLength: 6, Length: len(expected)}, + }, + }, + eventRecorder.Events(qlog.DatagramCreated{}), + ) + eventRecorder.Clear() + + data, err = serverConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, expected, data) +} + +func TestConnDatagramFailures(t *testing.T) { + t.Run("invalid varint", func(t *testing.T) { + testConnDatagramFailures(t, []byte{128}) + }) + + t.Run("invalid quarter stream ID", func(t *testing.T) { + testConnDatagramFailures(t, quicvarint.Append([]byte{}, maxQuarterStreamID+1)) + }) +} + +func testConnDatagramFailures(t *testing.T, datagram []byte) { + localConn, peerConn := newConnPair(t, withDatagrams()) + + conn := newRawConn(localConn, true, nil, nil, nil, nil) + + b := quicvarint.Append(nil, streamTypeControlStream) + b = (&settingsFrame{Datagram: true}).Append(b) + controlStr, err := peerConn.OpenUniStream() + require.NoError(t, err) + _, err = controlStr.Write(b) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + localStr, err := localConn.AcceptUniStream(ctx) + require.NoError(t, err) + + go conn.handleUnidirectionalStream(localStr, false) + + // Wait for SETTINGS to be received and datagram handling to start + select { + case <-conn.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for settings") + } + + require.NoError(t, peerConn.SendDatagram(datagram)) + + select { + case <-peerConn.Context().Done(): + require.ErrorIs(t, + context.Cause(peerConn.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeDatagramError)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout waiting for close") + } +} diff --git a/third_party/quic-go/http3/error.go b/third_party/quic-go/http3/error.go new file mode 100644 index 0000000..10f1e48 --- /dev/null +++ b/third_party/quic-go/http3/error.go @@ -0,0 +1,63 @@ +package http3 + +import ( + "errors" + "fmt" + + "github.com/apernet/quic-go" +) + +// Error is returned from the round tripper (for HTTP clients) +// and inside the HTTP handler (for HTTP servers) if an HTTP/3 error occurs. +// See section 8 of RFC 9114. +type Error struct { + Remote bool + ErrorCode ErrCode + ErrorMessage string +} + +var _ error = &Error{} + +func (e *Error) Error() string { + s := e.ErrorCode.string() + if s == "" { + s = fmt.Sprintf("H3 error (%#x)", uint64(e.ErrorCode)) + } + // Usually errors are remote. Only make it explicit for local errors. + if !e.Remote { + s += " (local)" + } + if e.ErrorMessage != "" { + s += ": " + e.ErrorMessage + } + return s +} + +func (e *Error) Is(target error) bool { + t, ok := target.(*Error) + return ok && e.ErrorCode == t.ErrorCode && e.Remote == t.Remote +} + +func maybeReplaceError(err error) error { + if err == nil { + return nil + } + + var ( + e Error + strErr *quic.StreamError + appErr *quic.ApplicationError + ) + switch { + default: + return err + case errors.As(err, &strErr): + e.Remote = strErr.Remote + e.ErrorCode = ErrCode(strErr.ErrorCode) + case errors.As(err, &appErr): + e.Remote = appErr.Remote + e.ErrorCode = ErrCode(appErr.ErrorCode) + e.ErrorMessage = appErr.ErrorMessage + } + return &e +} diff --git a/third_party/quic-go/http3/error_codes.go b/third_party/quic-go/http3/error_codes.go new file mode 100644 index 0000000..06918d0 --- /dev/null +++ b/third_party/quic-go/http3/error_codes.go @@ -0,0 +1,84 @@ +package http3 + +import ( + "fmt" + + "github.com/apernet/quic-go" +) + +type ErrCode quic.ApplicationErrorCode + +const ( + ErrCodeNoError ErrCode = 0x100 + ErrCodeGeneralProtocolError ErrCode = 0x101 + ErrCodeInternalError ErrCode = 0x102 + ErrCodeStreamCreationError ErrCode = 0x103 + ErrCodeClosedCriticalStream ErrCode = 0x104 + ErrCodeFrameUnexpected ErrCode = 0x105 + ErrCodeFrameError ErrCode = 0x106 + ErrCodeExcessiveLoad ErrCode = 0x107 + ErrCodeIDError ErrCode = 0x108 + ErrCodeSettingsError ErrCode = 0x109 + ErrCodeMissingSettings ErrCode = 0x10a + ErrCodeRequestRejected ErrCode = 0x10b + ErrCodeRequestCanceled ErrCode = 0x10c + ErrCodeRequestIncomplete ErrCode = 0x10d + ErrCodeMessageError ErrCode = 0x10e + ErrCodeConnectError ErrCode = 0x10f + ErrCodeVersionFallback ErrCode = 0x110 + ErrCodeDatagramError ErrCode = 0x33 + ErrCodeQPACKDecompressionFailed ErrCode = 0x200 +) + +func (e ErrCode) String() string { + s := e.string() + if s != "" { + return s + } + return fmt.Sprintf("unknown error code: %#x", uint16(e)) +} + +func (e ErrCode) string() string { + switch e { + case ErrCodeNoError: + return "H3_NO_ERROR" + case ErrCodeGeneralProtocolError: + return "H3_GENERAL_PROTOCOL_ERROR" + case ErrCodeInternalError: + return "H3_INTERNAL_ERROR" + case ErrCodeStreamCreationError: + return "H3_STREAM_CREATION_ERROR" + case ErrCodeClosedCriticalStream: + return "H3_CLOSED_CRITICAL_STREAM" + case ErrCodeFrameUnexpected: + return "H3_FRAME_UNEXPECTED" + case ErrCodeFrameError: + return "H3_FRAME_ERROR" + case ErrCodeExcessiveLoad: + return "H3_EXCESSIVE_LOAD" + case ErrCodeIDError: + return "H3_ID_ERROR" + case ErrCodeSettingsError: + return "H3_SETTINGS_ERROR" + case ErrCodeMissingSettings: + return "H3_MISSING_SETTINGS" + case ErrCodeRequestRejected: + return "H3_REQUEST_REJECTED" + case ErrCodeRequestCanceled: + return "H3_REQUEST_CANCELLED" + case ErrCodeRequestIncomplete: + return "H3_INCOMPLETE_REQUEST" + case ErrCodeMessageError: + return "H3_MESSAGE_ERROR" + case ErrCodeConnectError: + return "H3_CONNECT_ERROR" + case ErrCodeVersionFallback: + return "H3_VERSION_FALLBACK" + case ErrCodeDatagramError: + return "H3_DATAGRAM_ERROR" + case ErrCodeQPACKDecompressionFailed: + return "QPACK_DECOMPRESSION_FAILED" + default: + return "" + } +} diff --git a/third_party/quic-go/http3/error_codes_test.go b/third_party/quic-go/http3/error_codes_test.go new file mode 100644 index 0000000..af7642b --- /dev/null +++ b/third_party/quic-go/http3/error_codes_test.go @@ -0,0 +1,37 @@ +package http3 + +import ( + "go/ast" + "go/parser" + "go/token" + "path" + "runtime" + "strconv" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestErrorCodes(t *testing.T) { + // We parse the error code file, extract all constants, and verify that + // each of them has a string version. Go FTW! + _, thisfile, _, ok := runtime.Caller(0) + require.True(t, ok, "Failed to get current frame") + + filename := path.Join(path.Dir(thisfile), "error_codes.go") + fileAst, err := parser.ParseFile(token.NewFileSet(), filename, nil, 0) + require.NoError(t, err) + + constSpecs := fileAst.Decls[2].(*ast.GenDecl).Specs + require.Greater(t, len(constSpecs), 4) // at time of writing + + for _, c := range constSpecs { + valString := c.(*ast.ValueSpec).Values[0].(*ast.BasicLit).Value + val, err := strconv.ParseInt(valString, 0, 64) + require.NoError(t, err) + require.NotEqual(t, "unknown error code", ErrCode(val).String()) + } + + // Test unknown error code + require.Equal(t, "unknown error code: 0x1337", ErrCode(0x1337).String()) +} diff --git a/third_party/quic-go/http3/error_test.go b/third_party/quic-go/http3/error_test.go new file mode 100644 index 0000000..2b423fc --- /dev/null +++ b/third_party/quic-go/http3/error_test.go @@ -0,0 +1,87 @@ +package http3 + +import ( + "testing" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestErrorConversion(t *testing.T) { + tests := []struct { + name string + input error + expected error + }{ + {name: "nil error", input: nil, expected: nil}, + {name: "regular error", input: assert.AnError, expected: assert.AnError}, + { + name: "stream error", + input: &quic.StreamError{ErrorCode: 1337, Remote: true}, + expected: &Error{Remote: true, ErrorCode: 1337}, + }, + { + name: "application error", + input: &quic.ApplicationError{ErrorCode: 42, Remote: true, ErrorMessage: "foobar"}, + expected: &Error{Remote: true, ErrorCode: 42, ErrorMessage: "foobar"}, + }, + { + name: "transport error", + input: &quic.TransportError{ErrorCode: 42, Remote: true, ErrorMessage: "foobar"}, + expected: &quic.TransportError{ErrorCode: 42, Remote: true, ErrorMessage: "foobar"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := maybeReplaceError(tt.input) + if tt.expected == nil { + require.Nil(t, result) + } else { + require.ErrorIs(t, tt.expected, result) + } + }) + } +} + +func TestErrorString(t *testing.T) { + tests := []struct { + name string + err *Error + expected string + }{ + { + name: "remote error", + err: &Error{ErrorCode: 0x10c, Remote: true}, + expected: "H3_REQUEST_CANCELLED", + }, + { + name: "remote error with message", + err: &Error{ErrorCode: 0x10c, Remote: true, ErrorMessage: "foobar"}, + expected: "H3_REQUEST_CANCELLED: foobar", + }, + { + name: "local error", + err: &Error{ErrorCode: 0x10c, Remote: false}, + expected: "H3_REQUEST_CANCELLED (local)", + }, + { + name: "local error with message", + err: &Error{ErrorCode: 0x10c, Remote: false, ErrorMessage: "foobar"}, + expected: "H3_REQUEST_CANCELLED (local): foobar", + }, + { + name: "unknown error code", + err: &Error{ErrorCode: 0x1337, Remote: true}, + expected: "H3 error (0x1337)", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.expected, tt.err.Error()) + }) + } +} diff --git a/third_party/quic-go/http3/frames.go b/third_party/quic-go/http3/frames.go new file mode 100644 index 0000000..630b3dd --- /dev/null +++ b/third_party/quic-go/http3/frames.go @@ -0,0 +1,327 @@ +package http3 + +import ( + "bytes" + "errors" + "fmt" + "io" + "maps" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" +) + +// FrameType is the frame type of a HTTP/3 frame +type FrameType uint64 + +type frame any + +// The maximum length of an encoded HTTP/3 frame header is 16: +// The frame has a type and length field, both QUIC varints (maximum 8 bytes in length) +const frameHeaderLen = 16 + +type countingByteReader struct { + quicvarint.Reader + NumRead int +} + +func (r *countingByteReader) ReadByte() (byte, error) { + b, err := r.Reader.ReadByte() + if err == nil { + r.NumRead++ + } + return b, err +} + +func (r *countingByteReader) Read(b []byte) (int, error) { + n, err := r.Reader.Read(b) + r.NumRead += n + return n, err +} + +func (r *countingByteReader) Reset() { + r.NumRead = 0 +} + +type frameParser struct { + r io.Reader + streamID quic.StreamID + closeConn func(quic.ApplicationErrorCode, string) error +} + +func (p *frameParser) ParseNext(qlogger qlogwriter.Recorder) (frame, error) { + r := &countingByteReader{Reader: quicvarint.NewReader(p.r)} + for { + t, err := quicvarint.Read(r) + if err != nil { + return nil, err + } + l, err := quicvarint.Read(r) + if err != nil { + return nil, err + } + + switch t { + case 0x0: // DATA + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: p.streamID, + Raw: qlog.RawInfo{ + Length: int(l) + r.NumRead, + PayloadLength: int(l), + }, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }) + } + return &dataFrame{Length: l}, nil + case 0x1: // HEADERS + return &headersFrame{ + Length: l, + headerLen: r.NumRead, + }, nil + case 0x4: // SETTINGS + return parseSettingsFrame(r, l, p.streamID, qlogger) + case 0x3: // unsupported: CANCEL_PUSH + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: p.streamID, + Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.CancelPushFrame{}}, + }) + } + case 0x5: // unsupported: PUSH_PROMISE + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: p.streamID, + Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.PushPromiseFrame{}}, + }) + } + case 0x7: // GOAWAY + return parseGoAwayFrame(r, l, p.streamID, qlogger) + case 0xd: // unsupported: MAX_PUSH_ID + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: p.streamID, + Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.MaxPushIDFrame{}}, + }) + } + case 0x2, 0x6, 0x8, 0x9: // reserved frame types + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: p.streamID, + Raw: qlog.RawInfo{Length: r.NumRead + int(l), PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.ReservedFrame{Type: t}}, + }) + } + p.closeConn(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "") + return nil, fmt.Errorf("http3: reserved frame type: %d", t) + default: + // unknown frame types + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: p.streamID, + Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.UnknownFrame{Type: t}}, + }) + } + } + + // skip over the payload + if _, err := io.CopyN(io.Discard, r, int64(l)); err != nil { + return nil, err + } + r.Reset() + } +} + +type dataFrame struct { + Length uint64 +} + +func (f *dataFrame) Append(b []byte) []byte { + b = quicvarint.Append(b, 0x0) + return quicvarint.Append(b, f.Length) +} + +type headersFrame struct { + Length uint64 + headerLen int // number of bytes read for type and length field +} + +func (f *headersFrame) Append(b []byte) []byte { + b = quicvarint.Append(b, 0x1) + return quicvarint.Append(b, f.Length) +} + +const ( + // SETTINGS_MAX_FIELD_SECTION_SIZE + settingMaxFieldSectionSize = 0x6 + // Extended CONNECT, RFC 9220 + settingExtendedConnect = 0x8 + // HTTP Datagrams, RFC 9297 + settingDatagram = 0x33 +) + +type settingsFrame struct { + MaxFieldSectionSize int64 // SETTINGS_MAX_FIELD_SECTION_SIZE, -1 if not set + + Datagram bool // HTTP Datagrams, RFC 9297 + ExtendedConnect bool // Extended CONNECT, RFC 9220 + Other map[uint64]uint64 // all settings that we don't explicitly recognize +} + +func pointer[T any](v T) *T { + return &v +} + +func parseSettingsFrame(r *countingByteReader, l uint64, streamID quic.StreamID, qlogger qlogwriter.Recorder) (*settingsFrame, error) { + if l > 8*(1<<10) { + return nil, fmt.Errorf("unexpected size for SETTINGS frame: %d", l) + } + buf := make([]byte, l) + if _, err := io.ReadFull(r, buf); err != nil { + if err == io.ErrUnexpectedEOF { + return nil, io.EOF + } + return nil, err + } + frame := &settingsFrame{MaxFieldSectionSize: -1} + b := bytes.NewReader(buf) + settingsFrame := qlog.SettingsFrame{MaxFieldSectionSize: -1} + var readMaxFieldSectionSize, readDatagram, readExtendedConnect bool + for b.Len() > 0 { + id, err := quicvarint.Read(b) + if err != nil { // should not happen. We allocated the whole frame already. + return nil, err + } + val, err := quicvarint.Read(b) + if err != nil { // should not happen. We allocated the whole frame already. + return nil, err + } + + switch id { + case settingMaxFieldSectionSize: + if readMaxFieldSectionSize { + return nil, fmt.Errorf("duplicate setting: %d", id) + } + readMaxFieldSectionSize = true + frame.MaxFieldSectionSize = int64(val) + settingsFrame.MaxFieldSectionSize = int64(val) + case settingExtendedConnect: + if readExtendedConnect { + return nil, fmt.Errorf("duplicate setting: %d", id) + } + readExtendedConnect = true + if val != 0 && val != 1 { + return nil, fmt.Errorf("invalid value for SETTINGS_ENABLE_CONNECT_PROTOCOL: %d", val) + } + frame.ExtendedConnect = val == 1 + if qlogger != nil { + settingsFrame.ExtendedConnect = pointer(frame.ExtendedConnect) + } + case settingDatagram: + if readDatagram { + return nil, fmt.Errorf("duplicate setting: %d", id) + } + readDatagram = true + if val != 0 && val != 1 { + return nil, fmt.Errorf("invalid value for SETTINGS_H3_DATAGRAM: %d", val) + } + frame.Datagram = val == 1 + if qlogger != nil { + settingsFrame.Datagram = pointer(frame.Datagram) + } + default: + if _, ok := frame.Other[id]; ok { + return nil, fmt.Errorf("duplicate setting: %d", id) + } + if frame.Other == nil { + frame.Other = make(map[uint64]uint64) + } + frame.Other[id] = val + } + } + if qlogger != nil { + settingsFrame.Other = maps.Clone(frame.Other) + + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: streamID, + Raw: qlog.RawInfo{ + Length: r.NumRead, + PayloadLength: int(l), + }, + Frame: qlog.Frame{Frame: settingsFrame}, + }) + } + return frame, nil +} + +func (f *settingsFrame) Append(b []byte) []byte { + b = quicvarint.Append(b, 0x4) + var l int + if f.MaxFieldSectionSize >= 0 { + l += quicvarint.Len(settingMaxFieldSectionSize) + quicvarint.Len(uint64(f.MaxFieldSectionSize)) + } + for id, val := range f.Other { + l += quicvarint.Len(id) + quicvarint.Len(val) + } + if f.Datagram { + l += quicvarint.Len(settingDatagram) + quicvarint.Len(1) + } + if f.ExtendedConnect { + l += quicvarint.Len(settingExtendedConnect) + quicvarint.Len(1) + } + b = quicvarint.Append(b, uint64(l)) + if f.MaxFieldSectionSize >= 0 { + b = quicvarint.Append(b, settingMaxFieldSectionSize) + b = quicvarint.Append(b, uint64(f.MaxFieldSectionSize)) + } + if f.Datagram { + b = quicvarint.Append(b, settingDatagram) + b = quicvarint.Append(b, 1) + } + if f.ExtendedConnect { + b = quicvarint.Append(b, settingExtendedConnect) + b = quicvarint.Append(b, 1) + } + for id, val := range f.Other { + b = quicvarint.Append(b, id) + b = quicvarint.Append(b, val) + } + return b +} + +type goAwayFrame struct { + StreamID quic.StreamID +} + +func parseGoAwayFrame(r *countingByteReader, l uint64, streamID quic.StreamID, qlogger qlogwriter.Recorder) (*goAwayFrame, error) { + frame := &goAwayFrame{} + startLen := r.NumRead + id, err := quicvarint.Read(r) + if err != nil { + return nil, err + } + if r.NumRead-startLen != int(l) { + return nil, errors.New("GOAWAY frame: inconsistent length") + } + frame.StreamID = quic.StreamID(id) + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: streamID, + Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.GoAwayFrame{StreamID: frame.StreamID}}, + }) + } + return frame, nil +} + +func (f *goAwayFrame) Append(b []byte) []byte { + b = quicvarint.Append(b, 0x7) + b = quicvarint.Append(b, uint64(quicvarint.Len(uint64(f.StreamID)))) + return quicvarint.Append(b, uint64(f.StreamID)) +} diff --git a/third_party/quic-go/http3/frames_test.go b/third_party/quic-go/http3/frames_test.go new file mode 100644 index 0000000..c8faba2 --- /dev/null +++ b/third_party/quic-go/http3/frames_test.go @@ -0,0 +1,487 @@ +package http3 + +import ( + "bytes" + "context" + "fmt" + "io" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + "github.com/apernet/quic-go/testutils/events" + + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" + + "github.com/stretchr/testify/require" +) + +func testFrameParserEOF(t *testing.T, data []byte) { + t.Helper() + for i := range data { + b := make([]byte, i) + copy(b, data[:i]) + fp := frameParser{r: bytes.NewReader(b)} + _, err := fp.ParseNext(nil) + require.Error(t, err) + require.ErrorIs(t, err, io.EOF) + } +} + +func TestParserReservedFrameType(t *testing.T) { + for _, ft := range []uint64{0x2, 0x6, 0x8, 0x9} { + t.Run(fmt.Sprintf("type %#x", ft), func(t *testing.T) { + var eventRecorder events.Recorder + client, server := newConnPair(t, withDatagrams(), withServerRecorder(&eventRecorder)) + + data := quicvarint.Append(nil, ft) + data = quicvarint.Append(data, 6) + data = append(data, []byte("foobar")...) + + fp := frameParser{ + streamID: 42, + r: bytes.NewReader(data), + closeConn: client.CloseWithError, + } + _, err := fp.ParseNext(&eventRecorder) + require.Error(t, err) + require.ErrorContains(t, err, "http3: reserved frame type") + + select { + case <-server.Context().Done(): + require.ErrorIs(t, + context.Cause(server.Context()), + &quic.ApplicationError{Remote: true, ErrorCode: quic.ApplicationErrorCode(ErrCodeFrameUnexpected)}, + ) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameParsed{ + StreamID: 42, + Raw: qlog.RawInfo{Length: len(data), PayloadLength: 6}, + Frame: qlog.Frame{Frame: qlog.ReservedFrame{Type: ft}}, + }, + }, + eventRecorder.Events(qlog.FrameParsed{}), + ) + }) + } +} + +func TestParserUnknownFrameType(t *testing.T) { + data := quicvarint.Append(nil, 0xdead) + data = quicvarint.Append(data, 6) + data = append(data, []byte("foobar")...) + data = quicvarint.Append(data, 0xbeef) + data = quicvarint.Append(data, 3) + data = append(data, []byte("baz")...) + hf := &headersFrame{Length: 3} + data = hf.Append(data) + data = append(data, []byte("foo")...) + + r := bytes.NewReader(data) + fp := frameParser{r: r} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &headersFrame{}, f) + hf = f.(*headersFrame) + require.Equal(t, uint64(3), hf.Length) + payload := make([]byte, 3) + _, err = io.ReadFull(r, payload) + require.NoError(t, err) + require.Equal(t, []byte("foo"), payload) +} + +func TestParserUnsupportedFrameTypes(t *testing.T) { + for _, tc := range []struct { + name string + ft uint64 + qf any + }{ + {name: "CANCEL_PUSH", ft: 0x3, qf: qlog.CancelPushFrame{}}, + {name: "PUSH_PROMISE", ft: 0x5, qf: qlog.PushPromiseFrame{}}, + {name: "MAX_PUSH_ID", ft: 0xd, qf: qlog.MaxPushIDFrame{}}, + } { + t.Run(tc.name, func(t *testing.T) { + var eventRecorder events.Recorder + + data := quicvarint.Append(nil, tc.ft) + data = quicvarint.Append(data, 6) + data = append(data, []byte("foobar")...) + df := &dataFrame{Length: 3} + data = df.Append(data) + data = append(data, []byte("foo")...) + + r := bytes.NewReader(data) + fp := frameParser{streamID: 42, r: r} + + f, err := fp.ParseNext(&eventRecorder) + require.NoError(t, err) + require.IsType(t, &dataFrame{}, f) + df = f.(*dataFrame) + require.Equal(t, uint64(3), df.Length) + payload := make([]byte, 3) + _, err = io.ReadFull(r, payload) + require.NoError(t, err) + require.Equal(t, []byte("foo"), payload) + + headerLen := quicvarint.Len(tc.ft) + quicvarint.Len(6) + dfLen, _ := expectedFrameLength(t, df) + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameParsed{ + StreamID: 42, + Raw: qlog.RawInfo{Length: headerLen, PayloadLength: 6}, + Frame: qlog.Frame{Frame: tc.qf}, + }, + qlog.FrameParsed{ + StreamID: 42, + Raw: qlog.RawInfo{Length: dfLen, PayloadLength: 3}, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }, + }, + eventRecorder.Events(qlog.FrameParsed{}), + ) + }) + } +} + +func TestParserHeadersFrame(t *testing.T) { + data := quicvarint.Append(nil, 1) // type byte + data = quicvarint.Append(data, 0x1337) + fp := frameParser{r: bytes.NewReader(data)} + + // incomplete data results in an io.EOF + testFrameParserEOF(t, data) + + // parse + f1, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &headersFrame{}, f1) + require.Equal(t, uint64(0x1337), f1.(*headersFrame).Length) + + // write and parse + fp = frameParser{r: bytes.NewReader(f1.(*headersFrame).Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, f1, f2) +} + +func TestDataFrame(t *testing.T) { + data := quicvarint.Append(nil, 0) // type byte + data = quicvarint.Append(data, 0x1337) + fp := frameParser{r: bytes.NewReader(data)} + + // incomplete data results in an io.EOF + testFrameParserEOF(t, data) + + // parse + f1, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &dataFrame{}, f1) + require.Equal(t, uint64(0x1337), f1.(*dataFrame).Length) + + // write and parse + fp = frameParser{r: bytes.NewReader(f1.(*dataFrame).Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, f1, f2) +} + +func appendSetting(b []byte, key, value uint64) []byte { + b = quicvarint.Append(b, key) + b = quicvarint.Append(b, value) + return b +} + +func TestParserSettingsFrame(t *testing.T) { + settings := appendSetting(nil, 13, 37) + settings = appendSetting(settings, 0xdead, 0xbeef) + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + + // incomplete data results in an io.EOF + testFrameParserEOF(t, data) + + fp := frameParser{r: bytes.NewReader(data)} + frame, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, frame) + sf := frame.(*settingsFrame) + require.Len(t, sf.Other, 2) + require.Equal(t, uint64(37), sf.Other[uint64(13)]) + require.Equal(t, uint64(0xbeef), sf.Other[uint64(0xdead)]) + + // write and parse + fp = frameParser{r: bytes.NewReader(sf.Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f2) + sf2 := f2.(*settingsFrame) + require.Len(t, sf2.Other, len(sf.Other)) + require.Equal(t, sf.Other, sf2.Other) +} + +func TestParserSettingsFrameDuplicateSettings(t *testing.T) { + for _, tc := range []struct { + name string + num uint64 + val uint64 + }{ + { + name: "other setting", + num: 13, + val: 37, + }, + { + name: "extended connect", + num: settingExtendedConnect, + val: 1, + }, + { + name: "max field section size", + num: settingMaxFieldSectionSize, + val: 1337, + }, + { + name: "datagram", + num: settingDatagram, + val: 1, + }, + } { + t.Run(tc.name, func(t *testing.T) { + settings := appendSetting(nil, tc.num, tc.val) + settings = appendSetting(settings, tc.num, tc.val) + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + fp := frameParser{r: bytes.NewReader(data)} + _, err := fp.ParseNext(nil) + require.Error(t, err) + require.EqualError(t, err, fmt.Sprintf("duplicate setting: %d", tc.num)) + }) + } +} + +func TestParserSettingsFrameMaxFieldSectionSize(t *testing.T) { + t.Run("absent", func(t *testing.T) { + testParserSettingsFrameMaxFieldSectionSize(t, false) + }) + + t.Run("with value", func(t *testing.T) { + testParserSettingsFrameMaxFieldSectionSize(t, true) + }) +} + +func testParserSettingsFrameMaxFieldSectionSize(t *testing.T, present bool) { + var settings []byte + if present { + settings = appendSetting(nil, settingMaxFieldSectionSize, 1337) + } + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + + fp := frameParser{r: bytes.NewReader(data)} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f) + sf := f.(*settingsFrame) + if present { + require.EqualValues(t, 1337, sf.MaxFieldSectionSize) + } else { + require.EqualValues(t, -1, sf.MaxFieldSectionSize) + } + + fp = frameParser{r: bytes.NewReader(sf.Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, sf, f2) +} + +func TestParserSettingsFrameDatagram(t *testing.T) { + t.Run("enabled", func(t *testing.T) { + testParserSettingsFrameDatagram(t, true) + }) + t.Run("disabled", func(t *testing.T) { + testParserSettingsFrameDatagram(t, false) + }) +} + +func testParserSettingsFrameDatagram(t *testing.T, enabled bool) { + var settings []byte + switch enabled { + case true: + settings = appendSetting(nil, settingDatagram, 1) + case false: + settings = appendSetting(nil, settingDatagram, 0) + } + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + + fp := frameParser{r: bytes.NewReader(data)} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f) + sf := f.(*settingsFrame) + require.Equal(t, enabled, sf.Datagram) + + fp = frameParser{r: bytes.NewReader(sf.Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, sf, f2) +} + +func TestParserSettingsFrameDatagramInvalidValue(t *testing.T) { + settings := quicvarint.Append(nil, settingDatagram) + settings = quicvarint.Append(settings, 1337) + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + fp := frameParser{r: bytes.NewReader(data)} + _, err := fp.ParseNext(nil) + require.EqualError(t, err, "invalid value for SETTINGS_H3_DATAGRAM: 1337") +} + +func TestParserSettingsFrameExtendedConnect(t *testing.T) { + t.Run("enabled", func(t *testing.T) { + testParserSettingsFrameExtendedConnect(t, true) + }) + t.Run("disabled", func(t *testing.T) { + testParserSettingsFrameExtendedConnect(t, false) + }) +} + +func testParserSettingsFrameExtendedConnect(t *testing.T, enabled bool) { + var settings []byte + switch enabled { + case true: + settings = appendSetting(nil, settingExtendedConnect, 1) + case false: + settings = appendSetting(nil, settingExtendedConnect, 0) + } + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + + fp := frameParser{r: bytes.NewReader(data)} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f) + sf := f.(*settingsFrame) + require.Equal(t, enabled, sf.ExtendedConnect) + + fp = frameParser{r: bytes.NewReader(sf.Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, sf, f2) +} + +func TestParserSettingsFrameExtendedConnectInvalidValue(t *testing.T) { + settings := quicvarint.Append(nil, settingExtendedConnect) + settings = quicvarint.Append(settings, 1337) + data := quicvarint.Append(nil, 4) // type byte + data = quicvarint.Append(data, uint64(len(settings))) + data = append(data, settings...) + fp := frameParser{r: bytes.NewReader(data)} + _, err := fp.ParseNext(nil) + require.EqualError(t, err, "invalid value for SETTINGS_ENABLE_CONNECT_PROTOCOL: 1337") +} + +func TestParserGoAwayFrame(t *testing.T) { + data := quicvarint.Append(nil, 7) // type byte + data = quicvarint.Append(data, uint64(quicvarint.Len(100))) + data = quicvarint.Append(data, 100) + + // incomplete data results in an io.EOF + testFrameParserEOF(t, data) + + fp := frameParser{r: bytes.NewReader(data)} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &goAwayFrame{}, f) + require.Equal(t, quic.StreamID(100), f.(*goAwayFrame).StreamID) + + // write and parse + fp = frameParser{r: bytes.NewReader(f.(*goAwayFrame).Append(nil))} + f2, err := fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, f, f2) +} + +func FuzzFrameParser(f *testing.F) { + corpus := ossfuzzseeds.New(f) + + frames := []interface{ Append([]byte) []byte }{ + &dataFrame{Length: 5}, + &headersFrame{Length: 3}, + &settingsFrame{ + MaxFieldSectionSize: 1337, + Datagram: true, + ExtendedConnect: true, + Other: map[uint64]uint64{0xdead: 0xbeef}, + }, + &goAwayFrame{StreamID: 42}, + } + for _, fr := range frames { + corpus.Add(fr.Append(nil)) + } + + unknown := quicvarint.Append(nil, 0xdead) + unknown = quicvarint.Append(unknown, 6) + unknown = append(unknown, []byte("foobar")...) + corpus.Add(unknown) + + f.Fuzz(func(t *testing.T, data []byte) { + fp := frameParser{ + r: bytes.NewReader(data), + closeConn: func(quic.ApplicationErrorCode, string) error { return nil }, + } + for { + fr, err := fp.ParseNext(nil) + if err != nil { + return + } + + switch f := fr.(type) { + case *dataFrame: + if _, err := io.CopyN(io.Discard, fp.r, int64(f.Length)); err != nil { + return + } + case *headersFrame: + // Type and length are each at least one varint byte; HTTP/3 caps the pair at frameHeaderLen. + if f.headerLen < 2 || f.headerLen > frameHeaderLen { + t.Fatalf("HEADERS: headerLen %d outside [2, %d]", f.headerLen, frameHeaderLen) + } + if _, err := io.CopyN(io.Discard, fp.r, int64(f.Length)); err != nil { + return + } + case *settingsFrame: + // Unset uses -1; a present SETTINGS_MAX_FIELD_SECTION_SIZE is non-negative (see parseSettingsFrame). + if f.MaxFieldSectionSize != -1 && f.MaxFieldSectionSize < 0 { + t.Fatalf("SETTINGS: invalid MaxFieldSectionSize %d", f.MaxFieldSectionSize) + } + // Known settings are never stored in Other on a successful parse. + for id := range f.Other { + switch id { + case settingMaxFieldSectionSize, settingExtendedConnect, settingDatagram: + t.Fatalf("SETTINGS: known setting id %#x leaked into Other", id) + } + } + case *goAwayFrame: + // QUIC stream IDs fit in 62 bits; a negative value means uint64→int64 overflow in the parser. + if f.StreamID < 0 { + t.Fatalf("GOAWAY: negative StreamID %d", f.StreamID) + } + } + } + }) +} diff --git a/third_party/quic-go/http3/gzip_reader.go b/third_party/quic-go/http3/gzip_reader.go new file mode 100644 index 0000000..01983ac --- /dev/null +++ b/third_party/quic-go/http3/gzip_reader.go @@ -0,0 +1,39 @@ +package http3 + +// copied from net/transport.go + +// gzipReader wraps a response body so it can lazily +// call gzip.NewReader on the first call to Read +import ( + "compress/gzip" + "io" +) + +// call gzip.NewReader on the first call to Read +type gzipReader struct { + body io.ReadCloser // underlying Response.Body + zr *gzip.Reader // lazily-initialized gzip reader + zerr error // sticky error +} + +func newGzipReader(body io.ReadCloser) io.ReadCloser { + return &gzipReader{body: body} +} + +func (gz *gzipReader) Read(p []byte) (n int, err error) { + if gz.zerr != nil { + return 0, gz.zerr + } + if gz.zr == nil { + gz.zr, err = gzip.NewReader(gz.body) + if err != nil { + gz.zerr = err + return 0, err + } + } + return gz.zr.Read(p) +} + +func (gz *gzipReader) Close() error { + return gz.body.Close() +} diff --git a/third_party/quic-go/http3/headers.go b/third_party/quic-go/http3/headers.go new file mode 100644 index 0000000..9386386 --- /dev/null +++ b/third_party/quic-go/http3/headers.go @@ -0,0 +1,429 @@ +package http3 + +import ( + "bytes" + "errors" + "fmt" + "io" + "net/http" + "net/textproto" + "net/url" + "slices" + "strconv" + "strings" + + "golang.org/x/net/http/httpguts" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/quic-go/qpack" +) + +type qpackError struct{ err error } + +func (e *qpackError) Error() string { return fmt.Sprintf("qpack: %v", e.err) } +func (e *qpackError) Unwrap() error { return e.err } + +var errHeaderTooLarge = errors.New("http3: headers too large") + +type header struct { + // Pseudo header fields defined in RFC 9114 + Path string + Method string + Authority string + Scheme string + Status string + // for Extended connect + Protocol string + // parsed and deduplicated. -1 if no Content-Length header is sent + ContentLength int64 + // all non-pseudo headers + Headers http.Header +} + +// connection-specific header fields must not be sent on HTTP/3 +var invalidHeaderFields = [...]string{ + "connection", + "keep-alive", + "proxy-connection", + "transfer-encoding", + "upgrade", +} + +func parseHeaders(decodeFn qpack.DecodeFunc, isRequest bool, sizeLimit int, headerFields *[]qpack.HeaderField) (header, error) { + hdr := header{Headers: make(http.Header)} + var readFirstRegularHeader, readContentLength bool + var contentLengthStr string + for { + h, err := decodeFn() + if err != nil { + if err == io.EOF { + break + } + return header{}, &qpackError{err} + } + if headerFields != nil { + *headerFields = append(*headerFields, h) + } + // RFC 9114, section 4.2.2: + // The size of a field list is calculated based on the uncompressed size of fields, + // including the length of the name and value in bytes plus an overhead of 32 bytes for each field. + sizeLimit -= len(h.Name) + len(h.Value) + 32 + if sizeLimit < 0 { + return header{}, errHeaderTooLarge + } + if err := validateHeaderFieldNameAndValue(h); err != nil { + return header{}, err + } + if h.IsPseudo() { + if readFirstRegularHeader { + // all pseudo headers must appear before regular header fields, see section 4.3 of RFC 9114 + return header{}, fmt.Errorf("received pseudo header %s after a regular header field", h.Name) + } + var isResponsePseudoHeader bool // pseudo headers are either valid for requests or for responses + var isDuplicatePseudoHeader bool // pseudo headers are allowed to appear exactly once + switch h.Name { + case ":path": + isDuplicatePseudoHeader = hdr.Path != "" + hdr.Path = h.Value + case ":method": + isDuplicatePseudoHeader = hdr.Method != "" + hdr.Method = h.Value + case ":authority": + isDuplicatePseudoHeader = hdr.Authority != "" + hdr.Authority = h.Value + case ":protocol": // RFC 9220 + isDuplicatePseudoHeader = hdr.Protocol != "" + hdr.Protocol = h.Value + case ":scheme": + isDuplicatePseudoHeader = hdr.Scheme != "" + hdr.Scheme = h.Value + case ":status": + isDuplicatePseudoHeader = hdr.Status != "" + hdr.Status = h.Value + isResponsePseudoHeader = true + default: + return header{}, fmt.Errorf("unknown pseudo header: %s", h.Name) + } + if isDuplicatePseudoHeader { + return header{}, fmt.Errorf("duplicate pseudo header: %s", h.Name) + } + if isRequest && isResponsePseudoHeader { + return header{}, fmt.Errorf("invalid request pseudo header: %s", h.Name) + } + if !isRequest && !isResponsePseudoHeader { + return header{}, fmt.Errorf("invalid response pseudo header: %s", h.Name) + } + } else { + if err := validateRegularHeaderField(h); err != nil { + return header{}, err + } + readFirstRegularHeader = true + switch h.Name { + case "content-length": + // Ignore duplicate Content-Length headers. + // Fail if the duplicates differ. + if !readContentLength { + readContentLength = true + contentLengthStr = h.Value + } else if contentLengthStr != h.Value { + return header{}, fmt.Errorf("contradicting content lengths (%s and %s)", contentLengthStr, h.Value) + } + default: + hdr.Headers.Add(h.Name, h.Value) + } + } + } + hdr.ContentLength = -1 + if len(contentLengthStr) > 0 { + // use ParseUint instead of ParseInt, so that parsing fails on negative values + cl, err := strconv.ParseUint(contentLengthStr, 10, 63) + if err != nil { + return header{}, fmt.Errorf("invalid content length: %w", err) + } + hdr.Headers.Set("Content-Length", contentLengthStr) + hdr.ContentLength = int64(cl) + } + return hdr, nil +} + +func validateHeaderFieldNameAndValue(h qpack.HeaderField) error { + // field names need to be lowercase, see section 4.2 of RFC 9114 + if strings.ToLower(h.Name) != h.Name { + return fmt.Errorf("header field is not lower-case: %s", h.Name) + } + if !httpguts.ValidHeaderFieldValue(h.Value) { + return fmt.Errorf("invalid header field value for %s: %q", h.Name, h.Value) + } + return nil +} + +func validateRegularHeaderField(h qpack.HeaderField) error { + if !httpguts.ValidHeaderFieldName(h.Name) { + return fmt.Errorf("invalid header field name: %q", h.Name) + } + if slices.Contains(invalidHeaderFields[:], h.Name) { + return fmt.Errorf("invalid header field name: %q", h.Name) + } + if h.Name == "te" && h.Value != "trailers" { + return fmt.Errorf("invalid TE header field value: %q", h.Value) + } + return nil +} + +func validateTrailerHeaderField(h qpack.HeaderField) error { + if err := validateRegularHeaderField(h); err != nil { + return err + } + if !httpguts.ValidTrailerHeader(h.Name) { + return fmt.Errorf("invalid trailer field name: %q", h.Name) + } + return nil +} + +func parseTrailers(decodeFn qpack.DecodeFunc, sizeLimit int, headerFields *[]qpack.HeaderField) (http.Header, error) { + h := make(http.Header) + for { + hf, err := decodeFn() + if err != nil { + if err == io.EOF { + break + } + return nil, &qpackError{err} + } + if headerFields != nil { + *headerFields = append(*headerFields, hf) + } + // RFC 9114, section 4.2.2: + // The size of a field list is calculated based on the uncompressed size of fields, + // including the length of the name and value in bytes plus an overhead of 32 bytes for each field. + sizeLimit -= len(hf.Name) + len(hf.Value) + 32 + if sizeLimit < 0 { + return nil, errHeaderTooLarge + } + if err := validateHeaderFieldNameAndValue(hf); err != nil { + return nil, err + } + if hf.IsPseudo() { + return nil, fmt.Errorf("http3: received pseudo header in trailer: %s", hf.Name) + } + if err := validateTrailerHeaderField(hf); err != nil { + return nil, err + } + h.Add(hf.Name, hf.Value) + } + return h, nil +} + +func requestFromHeaders(decodeFn qpack.DecodeFunc, sizeLimit int, headerFields *[]qpack.HeaderField) (*http.Request, error) { + hdr, err := parseHeaders(decodeFn, true, sizeLimit, headerFields) + if err != nil { + return nil, err + } + // concatenate cookie headers, see https://tools.ietf.org/html/rfc6265#section-5.4 + if len(hdr.Headers["Cookie"]) > 0 { + hdr.Headers.Set("Cookie", strings.Join(hdr.Headers["Cookie"], "; ")) + } + + isConnect := hdr.Method == http.MethodConnect + // Extended CONNECT, see https://datatracker.ietf.org/doc/html/rfc8441#section-4 + isExtendedConnected := isConnect && hdr.Protocol != "" + if isExtendedConnected { + if !validExtendedConnectProtocol(hdr.Protocol) { + return nil, fmt.Errorf("invalid :protocol: %q", hdr.Protocol) + } + if hdr.Scheme == "" || hdr.Path == "" || hdr.Authority == "" { + return nil, errors.New("extended CONNECT: :scheme, :path and :authority must not be empty") + } + } else if isConnect { + if hdr.Path != "" || hdr.Authority == "" { // normal CONNECT + return nil, errors.New(":path must be empty and :authority must not be empty") + } + } else if len(hdr.Path) == 0 || len(hdr.Authority) == 0 || len(hdr.Method) == 0 { + return nil, errors.New(":path, :authority and :method must not be empty") + } + + if !isExtendedConnected && len(hdr.Protocol) > 0 { + return nil, errors.New(":protocol must be empty") + } + + var u *url.URL + var requestURI string + + protocol := "HTTP/3.0" + + if isConnect { + u = &url.URL{} + if isExtendedConnected { + u, err = url.ParseRequestURI(hdr.Path) + if err != nil { + return nil, err + } + protocol = hdr.Protocol + } else { + u.Path = hdr.Path + } + requestURI = hdr.Authority + } else { + u, err = url.ParseRequestURI(hdr.Path) + if err != nil { + return nil, fmt.Errorf("invalid request URI: %w", err) + } + requestURI = hdr.Path + } + u.Scheme = hdr.Scheme + u.Host = hdr.Authority + + req := &http.Request{ + Method: hdr.Method, + URL: u, + Proto: protocol, + ProtoMajor: 3, + ProtoMinor: 0, + Header: hdr.Headers, + Body: nil, + ContentLength: hdr.ContentLength, + Host: hdr.Authority, + RequestURI: requestURI, + } + req.Trailer = extractAnnouncedTrailers(req.Header) + return req, nil +} + +func validExtendedConnectProtocol(protocol string) bool { + // RFC 9220 specifies that the semantics of the :protocol pseudo are the same as defined in RFC 8441. + // RFC 8441, Section 4 specifies that :protocol is a single value from the HTTP Upgrade Token Registry. + // RFC 9110, Section 16.7 specifies that HTTP Upgrade Token Registry uses token grammar. + // Therefore, ValidHeaderFieldName is the right syntax check here, despite the misleading name. + return httpguts.ValidHeaderFieldName(protocol) +} + +// updateResponseFromHeaders sets up http.Response as an HTTP/3 response, +// using the decoded qpack header filed. +// It is only called for the HTTP header (and not the HTTP trailer). +// It takes an http.Response as an argument to allow the caller to set the trailer later on. +func updateResponseFromHeaders(rsp *http.Response, decodeFn qpack.DecodeFunc, sizeLimit int, headerFields *[]qpack.HeaderField) error { + hdr, err := parseHeaders(decodeFn, false, sizeLimit, headerFields) + if err != nil { + return err + } + if hdr.Status == "" { + return errors.New("missing :status field") + } + rsp.Proto = "HTTP/3.0" + rsp.ProtoMajor = 3 + rsp.Header = hdr.Headers + rsp.Trailer = extractAnnouncedTrailers(rsp.Header) + rsp.ContentLength = hdr.ContentLength + + status, err := strconv.Atoi(hdr.Status) + if err != nil { + return fmt.Errorf("invalid status code: %w", err) + } + rsp.StatusCode = status + rsp.Status = hdr.Status + " " + http.StatusText(status) + return nil +} + +// extractAnnouncedTrailers extracts trailer keys from the "Trailer" header. +// It returns a map with the announced keys set to nil values, and removes the "Trailer" header. +// It handles both duplicate as well as comma-separated values for the Trailer header. +// For example: +// +// Trailer: Trailer1, Trailer2 +// Trailer: Trailer3 +// +// Will result in a map containing the keys "Trailer1", "Trailer2", "Trailer3" with nil values. +func extractAnnouncedTrailers(header http.Header) http.Header { + rawTrailers, ok := header["Trailer"] + if !ok { + return nil + } + + trailers := make(http.Header) + for _, rawVal := range rawTrailers { + for val := range strings.SplitSeq(rawVal, ",") { + trailers[http.CanonicalHeaderKey(textproto.TrimString(val))] = nil + } + } + delete(header, "Trailer") + return trailers +} + +// writeTrailers encodes and writes HTTP trailers as a HEADERS frame. +// It returns true if trailers were written, false if there were no trailers to write. +func writeTrailers(wr io.Writer, trailers http.Header, streamID quic.StreamID, qlogger qlogwriter.Recorder) (bool, error) { + var hasValues bool + for k, vals := range trailers { + if httpguts.ValidTrailerHeader(k) && len(vals) > 0 { + hasValues = true + break + } + } + if !hasValues { + return false, nil + } + + var buf bytes.Buffer + enc := qpack.NewEncoder(&buf) + var headerFields []qlog.HeaderField + if qlogger != nil { + headerFields = make([]qlog.HeaderField, 0, len(trailers)) + } + + for k, vals := range trailers { + if len(vals) == 0 { + continue + } + if !httpguts.ValidTrailerHeader(k) { + continue + } + lowercaseKey := strings.ToLower(k) + for _, v := range vals { + if err := enc.WriteField(qpack.HeaderField{Name: lowercaseKey, Value: v}); err != nil { + return false, err + } + if qlogger != nil { + headerFields = append(headerFields, qlog.HeaderField{Name: lowercaseKey, Value: v}) + } + } + } + + b := make([]byte, 0, frameHeaderLen+buf.Len()) + b = (&headersFrame{Length: uint64(buf.Len())}).Append(b) + b = append(b, buf.Bytes()...) + if qlogger != nil { + qlogCreatedHeadersFrame(qlogger, streamID, len(b), buf.Len(), headerFields) + } + _, err := wr.Write(b) + return true, err +} + +func decodeTrailers(r io.Reader, hf *headersFrame, maxHeaderBytes int, decoder *qpack.Decoder, qlogger qlogwriter.Recorder, streamID quic.StreamID) (http.Header, error) { + if hf.Length > uint64(maxHeaderBytes) { + maybeQlogInvalidHeadersFrame(qlogger, streamID, hf.Length) + return nil, fmt.Errorf("http3: HEADERS frame too large: %d bytes (max: %d)", hf.Length, maxHeaderBytes) + } + + b := make([]byte, hf.Length) + if _, err := io.ReadFull(r, b); err != nil { + return nil, err + } + decodeFn := decoder.Decode(b) + var fields []qpack.HeaderField + var headerFields *[]qpack.HeaderField + if qlogger != nil { + fields = make([]qpack.HeaderField, 0, 16) + headerFields = &fields + } + trailers, err := parseTrailers(decodeFn, maxHeaderBytes, headerFields) + if err != nil { + maybeQlogInvalidHeadersFrame(qlogger, streamID, hf.Length) + return nil, err + } + if qlogger != nil { + qlogParsedHeadersFrame(qlogger, streamID, hf, fields) + } + return trailers, nil +} diff --git a/third_party/quic-go/http3/headers_test.go b/third_party/quic-go/http3/headers_test.go new file mode 100644 index 0000000..48070dc --- /dev/null +++ b/third_party/quic-go/http3/headers_test.go @@ -0,0 +1,776 @@ +package http3 + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "math" + "net/http" + "testing" + + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" + "github.com/quic-go/qpack" + + "github.com/stretchr/testify/require" + "golang.org/x/net/http/httpguts" +) + +func decodeFromSlice(headers []qpack.HeaderField) qpack.DecodeFunc { + var i int + return func() (qpack.HeaderField, error) { + if i >= len(headers) { + return qpack.HeaderField{}, io.EOF + } + h := headers[i] + i++ + return h, nil + } +} + +func TestRequestHeaderParsing(t *testing.T) { + t.Run("regular path", func(t *testing.T) { + testRequestHeaderParsing(t, "/foo") + }) + + // see https://github.com/apernet/quic-go/pull/1898 + t.Run("path starting with //", func(t *testing.T) { + testRequestHeaderParsing(t, "//foo") + }) +} + +func testRequestHeaderParsing(t *testing.T, path string) { + headers := []qpack.HeaderField{ + {Name: ":scheme", Value: "https"}, + {Name: ":path", Value: path}, + {Name: ":authority", Value: "quic-go.net:443"}, + {Name: ":method", Value: http.MethodGet}, + {Name: "content-length", Value: "42"}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, http.MethodGet, req.Method) + require.Equal(t, path, req.URL.Path) + require.Equal(t, "quic-go.net:443", req.URL.Host) + require.Equal(t, "HTTP/3.0", req.Proto) + require.Equal(t, 3, req.ProtoMajor) + require.Zero(t, req.ProtoMinor) + require.Equal(t, int64(42), req.ContentLength) + require.Equal(t, 1, len(req.Header)) + require.Equal(t, "42", req.Header.Get("Content-Length")) + require.Nil(t, req.Body) + require.Equal(t, "quic-go.net:443", req.Host) + require.Equal(t, path, req.RequestURI) + require.Equal(t, "quic-go.net", req.URL.Hostname()) + require.Equal(t, "https", req.URL.Scheme) + require.Equal(t, "443", req.URL.Port()) +} + +func TestRequestHeadersContentLength(t *testing.T) { + t.Run("no content length", func(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":path", Value: "/"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, int64(-1), req.ContentLength) + }) + + t.Run("multiple content lengths", func(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":path", Value: "/"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + {Name: "content-length", Value: "42"}, + {Name: "content-length", Value: "42"}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, "42", req.Header.Get("Content-Length")) + }) +} + +func TestRequestHeadersContentLengthValidation(t *testing.T) { + for _, tc := range []struct { + name string + headers []qpack.HeaderField + err string + errContains string + }{ + { + name: "negative content length", + headers: []qpack.HeaderField{ + {Name: "content-length", Value: "-42"}, + }, + errContains: "invalid content length", + }, + { + name: "multiple differing content lengths", + headers: []qpack.HeaderField{ + {Name: "content-length", Value: "42"}, + {Name: "content-length", Value: "1337"}, + }, + err: "contradicting content lengths (42 and 1337)", + }, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := requestFromHeaders(decodeFromSlice(tc.headers), math.MaxInt, nil) + if tc.errContains != "" { + require.ErrorContains(t, err, tc.errContains) + } + if tc.err != "" { + require.EqualError(t, err, tc.err) + } + }) + } +} + +func TestRequestHeadersValidation(t *testing.T) { + for _, tc := range []struct { + name string + headers []qpack.HeaderField + err string + errContains string + }{ + { + name: "upper-case field name", + headers: []qpack.HeaderField{ + {Name: "Content-Length", Value: "42"}, + }, + err: "header field is not lower-case: Content-Length", + }, + { + name: "unknown pseudo header", + headers: []qpack.HeaderField{ + {Name: ":foo", Value: "bar"}, + }, + err: "unknown pseudo header: :foo", + }, + { + name: "pseudo header after regular header", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: "content-length", Value: "42"}, + {Name: ":authority", Value: "quic-go.net"}, + }, + err: "received pseudo header :authority after a regular header field", + }, + { + name: "invalid field name", + headers: []qpack.HeaderField{ + {Name: "@", Value: "42"}, + }, + err: `invalid header field name: "@"`, + }, + { + name: "invalid field value", + headers: []qpack.HeaderField{ + {Name: "content", Value: "\n"}, + }, + err: `invalid header field value for content: "\n"`, + }, + { + name: ":status header field", // :status is a response pseudo header + headers: []qpack.HeaderField{ + {Name: ":status", Value: "404"}, + }, + err: "invalid request pseudo header: :status", + }, + { + name: "missing :path", + headers: []qpack.HeaderField{ + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + }, + err: ":path, :authority and :method must not be empty", + }, + { + name: "missing :authority", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":method", Value: http.MethodGet}, + }, + err: ":path, :authority and :method must not be empty", + }, + { + name: "missing :method", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":authority", Value: "quic-go.net"}, + }, + err: ":path, :authority and :method must not be empty", + }, + { + name: "duplicate :path", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":path", Value: "/foo"}, + }, + err: "duplicate pseudo header: :path", + }, + { + name: "duplicate :authority", + headers: []qpack.HeaderField{ + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":authority", Value: "quic-go.net"}, + }, + err: "duplicate pseudo header: :authority", + }, + { + name: "duplicate :method", + headers: []qpack.HeaderField{ + {Name: ":method", Value: http.MethodGet}, + {Name: ":method", Value: http.MethodGet}, + }, + err: "duplicate pseudo header: :method", + }, + { + name: "invalid :protocol", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + {Name: ":protocol", Value: "connect-udp"}, + }, + err: ":protocol must be empty", + }, + { + name: "invalid :path", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "invalid path"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + }, + errContains: "invalid request URI", + }, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := requestFromHeaders(decodeFromSlice(tc.headers), math.MaxInt, nil) + if tc.errContains != "" { + require.ErrorContains(t, err, tc.errContains) + } + if tc.err != "" { + require.EqualError(t, err, tc.err) + } + require.NotErrorAs(t, err, new(*qpackError)) + }) + } +} + +func TestCookieHeader(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + {Name: "cookie", Value: "cookie1=foobar1"}, + {Name: "cookie", Value: "cookie2=foobar2"}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, http.Header{ + "Cookie": []string{"cookie1=foobar1; cookie2=foobar2"}, + }, req.Header) +} + +func TestHeadersConcatenation(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodGet}, + {Name: "cache-control", Value: "max-age=0"}, + {Name: "duplicate-header", Value: "1"}, + {Name: "duplicate-header", Value: "2"}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, http.Header{ + "Cache-Control": []string{"max-age=0"}, + "Duplicate-Header": []string{"1", "2"}, + }, req.Header) +} + +func TestRequestHeadersConnect(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodConnect}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, http.MethodConnect, req.Method) + require.Equal(t, "HTTP/3.0", req.Proto) + require.Equal(t, "quic-go.net", req.RequestURI) +} + +func TestRequestHeadersConnectValidation(t *testing.T) { + for _, tc := range []struct { + name string + headers []qpack.HeaderField + err string + }{ + { + name: "missing :authority", + headers: []qpack.HeaderField{ + {Name: ":method", Value: http.MethodConnect}, + }, + err: ":path must be empty and :authority must not be empty", + }, + { + name: ":path set", + headers: []qpack.HeaderField{ + {Name: ":path", Value: "/foo"}, + {Name: ":method", Value: http.MethodConnect}, + }, + err: ":path must be empty and :authority must not be empty", + }, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := requestFromHeaders(decodeFromSlice(tc.headers), math.MaxInt, nil) + require.EqualError(t, err, tc.err) + }) + } +} + +func TestRequestHeadersExtendedConnect(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":protocol", Value: "webtransport"}, + {Name: ":scheme", Value: "ftp"}, + {Name: ":method", Value: http.MethodConnect}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":path", Value: "/foo?val=1337"}, + } + req, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, http.MethodConnect, req.Method) + require.Equal(t, "webtransport", req.Proto) + require.Equal(t, "ftp://quic-go.net/foo?val=1337", req.URL.String()) + require.Equal(t, "1337", req.URL.Query().Get("val")) +} + +func TestRequestHeadersExtendedConnectRequestValidation(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":protocol", Value: "webtransport"}, + {Name: ":method", Value: http.MethodConnect}, + {Name: ":authority", Value: "quic.clemente.io"}, + {Name: ":path", Value: "/foo"}, + } + _, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.EqualError(t, err, "extended CONNECT: :scheme, :path and :authority must not be empty") +} + +func TestRequestHeadersExtendedConnectInvalidProtocol(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":protocol", Value: "HTTP/3.0"}, + {Name: ":scheme", Value: "https"}, + {Name: ":method", Value: http.MethodConnect}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":path", Value: "/foo"}, + } + _, err := requestFromHeaders(decodeFromSlice(headers), math.MaxInt, nil) + require.EqualError(t, err, `invalid :protocol: "HTTP/3.0"`) +} + +func TestResponseHeaderParsing(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":status", Value: "200"}, + {Name: "content-length", Value: "42"}, + } + rsp := &http.Response{} + require.NoError(t, updateResponseFromHeaders(rsp, decodeFromSlice(headers), math.MaxInt, nil)) + require.Equal(t, "HTTP/3.0", rsp.Proto) + require.Equal(t, 3, rsp.ProtoMajor) + require.Zero(t, rsp.ProtoMinor) + require.Equal(t, int64(42), rsp.ContentLength) + require.Equal(t, 1, len(rsp.Header)) + require.Equal(t, "42", rsp.Header.Get("Content-Length")) + require.Nil(t, rsp.Body) + require.Equal(t, 200, rsp.StatusCode) + require.Equal(t, "200 OK", rsp.Status) +} + +func TestResponseHeaderParsingValidation(t *testing.T) { + for _, tc := range []struct { + name string + headers []qpack.HeaderField + err string + errContains string + }{ + { + name: "missing :status", + headers: []qpack.HeaderField{ + {Name: "content-length", Value: "42"}, + }, + err: "missing :status field", + }, + { + name: "invalid status code", + headers: []qpack.HeaderField{ + {Name: ":status", Value: "foobar"}, + }, + errContains: "invalid status code", + }, + { + name: ":method header field", // :method is a request pseudo header + headers: []qpack.HeaderField{ + {Name: ":method", Value: http.MethodGet}, + }, + err: "invalid response pseudo header: :method", + }, + { + name: "duplicate :status", + headers: []qpack.HeaderField{ + {Name: ":status", Value: "200"}, + {Name: ":status", Value: "404"}, + }, + err: "duplicate pseudo header: :status", + }, + } { + t.Run(tc.name, func(t *testing.T) { + err := updateResponseFromHeaders(&http.Response{}, decodeFromSlice(tc.headers), math.MaxInt, nil) + if tc.errContains != "" { + require.ErrorContains(t, err, tc.errContains) + } + if tc.err != "" { + require.EqualError(t, err, tc.err) + } + }) + } + + for _, tc := range []struct { + name string + invalidField string + }{ + {name: "connection", invalidField: "connection"}, + {name: "keep-alive", invalidField: "keep-alive"}, + {name: "proxy-connection", invalidField: "proxy-connection"}, + {name: "transfer-encoding", invalidField: "transfer-encoding"}, + {name: "upgrade", invalidField: "upgrade"}, + } { + t.Run("invalid field: "+tc.name, func(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":status", Value: "404"}, + {Name: tc.invalidField, Value: "some-value"}, + } + err := updateResponseFromHeaders(&http.Response{}, decodeFromSlice(headers), math.MaxInt, nil) + require.EqualError(t, err, fmt.Sprintf("invalid header field name: %q", tc.invalidField)) + }) + } +} + +func TestResponseTrailerFields(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":status", Value: "200"}, + {Name: "trailer", Value: "Trailer1, Trailer2"}, + {Name: "trailer", Value: "TRAILER3"}, + } + var rsp http.Response + require.NoError(t, updateResponseFromHeaders(&rsp, decodeFromSlice(headers), math.MaxInt, nil)) + require.Equal(t, 0, len(rsp.Header)) + require.Equal(t, http.Header(map[string][]string{ + "Trailer1": nil, + "Trailer2": nil, + "Trailer3": nil, + }), rsp.Trailer) +} + +func TestResponseTrailerParsingTE(t *testing.T) { + headers := []qpack.HeaderField{ + {Name: ":status", Value: "404"}, + {Name: "te", Value: "trailers"}, + } + require.NoError(t, updateResponseFromHeaders(&http.Response{}, decodeFromSlice(headers), math.MaxInt, nil)) + headers = []qpack.HeaderField{ + {Name: ":status", Value: "404"}, + {Name: "te", Value: "not-trailers"}, + } + require.EqualError(t, + updateResponseFromHeaders(&http.Response{}, decodeFromSlice(headers), math.MaxInt, nil), + `invalid TE header field value: "not-trailers"`) +} + +func TestResponseTrailerParsing(t *testing.T) { + trailerHdr, err := parseTrailers(decodeFromSlice([]qpack.HeaderField{ + {Name: "foo", Value: "42"}, + }), math.MaxInt, nil) + require.NoError(t, err) + require.Equal(t, "42", trailerHdr.Get("Foo")) +} + +func TestResponseTrailerParsingValidation(t *testing.T) { + for _, tc := range []struct { + name string + headers []qpack.HeaderField + sizeLimit int + err string + errContains string + errIs error + }{ + { + name: "field list too large", + headers: []qpack.HeaderField{ + {Name: "foo", Value: "bar"}, + }, + sizeLimit: 5, + errIs: errHeaderTooLarge, + }, + { + name: "upper-case field name", + headers: []qpack.HeaderField{ + {Name: "Foo", Value: "bar"}, + }, + err: "header field is not lower-case: Foo", + }, + { + name: "pseudo header", + headers: []qpack.HeaderField{ + {Name: ":status", Value: "200"}, + }, + err: "http3: received pseudo header in trailer: :status", + }, + { + name: "invalid field name", + headers: []qpack.HeaderField{ + {Name: "@", Value: "bar"}, + }, + err: `invalid header field name: "@"`, + }, + { + name: "invalid field value", + headers: []qpack.HeaderField{ + {Name: "foo", Value: "\n"}, + }, + err: `invalid header field value for foo: "\n"`, + }, + { + name: "connection-specific field", + headers: []qpack.HeaderField{ + {Name: "connection", Value: "close"}, + }, + err: `invalid header field name: "connection"`, + }, + { + name: "invalid te field value", + headers: []qpack.HeaderField{ + {Name: "te", Value: "gzip"}, + }, + err: `invalid TE header field value: "gzip"`, + }, + { + name: "invalid trailer field", + headers: []qpack.HeaderField{ + {Name: "content-length", Value: "42"}, + }, + err: `invalid trailer field name: "content-length"`, + }, + { + name: "valid header field name disallowed in trailers", + headers: []qpack.HeaderField{ + {Name: "if-match", Value: "etag"}, + }, + err: `invalid trailer field name: "if-match"`, + }, + } { + t.Run(tc.name, func(t *testing.T) { + sizeLimit := tc.sizeLimit + if sizeLimit == 0 { + sizeLimit = math.MaxInt + } + _, err := parseTrailers(decodeFromSlice(tc.headers), sizeLimit, nil) + if tc.errIs != nil { + require.ErrorIs(t, err, tc.errIs) + } + if tc.errContains != "" { + require.ErrorContains(t, err, tc.errContains) + } + if tc.err != "" { + require.EqualError(t, err, tc.err) + } + require.NotErrorAs(t, err, new(*qpackError)) + }) + } +} + +func TestQpackError(t *testing.T) { + buf := &bytes.Buffer{} + enc := qpack.NewEncoder(buf) + enc.WriteField(qpack.HeaderField{Name: ":status", Value: "200"}) + enc.Close() + + t.Run("header parsing", func(t *testing.T) { + dec := qpack.NewDecoder() + decodeFn := dec.Decode(buf.Bytes()[:len(buf.Bytes())/2]) + _, err := requestFromHeaders(decodeFn, math.MaxInt, nil) + require.ErrorAs(t, err, new(*qpackError)) + }) + + t.Run("trailer parsing", func(t *testing.T) { + dec := qpack.NewDecoder() + decodeFn := dec.Decode(buf.Bytes()[:len(buf.Bytes())/2]) + err := updateResponseFromHeaders(&http.Response{}, decodeFn, math.MaxInt, nil) + require.ErrorAs(t, err, new(*qpackError)) + }) +} + +func BenchmarkRequestFromHeaders(b *testing.B) { + b.ReportAllocs() + + headers := []qpack.HeaderField{ + {Name: ":path", Value: "/api/v1/users/12345"}, + {Name: ":authority", Value: "quic-go.net"}, + {Name: ":method", Value: http.MethodPost}, + {Name: "content-type", Value: "application/json"}, + {Name: "content-length", Value: "1024"}, + {Name: "user-agent", Value: "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/26.0 Safari/605.1.15"}, + {Name: "accept", Value: "application/json, text/plain, */*"}, + {Name: "accept-encoding", Value: "gzip, deflate, br"}, + {Name: "accept-language", Value: "en-US,en;q=0.9"}, + {Name: "cache-control", Value: "no-cache"}, + {Name: "cookie", Value: "session_id=abc123"}, + {Name: "cookie", Value: "user_pref=dark_mode"}, + {Name: "referer", Value: "https://quic-go.net/docs/http3/"}, + } + var buf bytes.Buffer + enc := qpack.NewEncoder(&buf) + for _, hf := range headers { + require.NoError(b, enc.WriteField(hf)) + } + + dec := qpack.NewDecoder() + for b.Loop() { + decodeFn := dec.Decode(buf.Bytes()) + if _, err := requestFromHeaders(decodeFn, math.MaxInt, nil); err != nil { + b.Fatalf("failed to parse request: %v", err) + } + } +} + +func FuzzHeaderParsing(f *testing.F) { + corpus := ossfuzzseeds.New(f) + + for _, s := range [][]qpack.HeaderField{ + { // GET request + {Name: ":method", Value: "GET"}, + {Name: ":scheme", Value: "https"}, + {Name: ":path", Value: "/"}, + {Name: ":authority", Value: "example.com"}, + }, + { // POST with Content-Length + {Name: ":method", Value: "POST"}, + {Name: ":scheme", Value: "https"}, + {Name: ":path", Value: "/submit"}, + {Name: ":authority", Value: "example.com"}, + {Name: "content-length", Value: "42"}, + {Name: "content-type", Value: "application/json"}, + }, + { // CONNECT request + {Name: ":method", Value: "CONNECT"}, + {Name: ":authority", Value: "proxy.example.com:443"}, + }, + { // extended CONNECT + {Name: ":method", Value: "CONNECT"}, + {Name: ":scheme", Value: "https"}, + {Name: ":path", Value: "/webtransport"}, + {Name: ":authority", Value: "example.com"}, + {Name: ":protocol", Value: "webtransport"}, + }, + { // 200 response + {Name: ":status", Value: "200"}, + {Name: "content-type", Value: "text/html"}, + {Name: "content-length", Value: "1024"}, + }, + { // response with trailer announcement + {Name: ":status", Value: "200"}, + {Name: "trailer", Value: "Checksum"}, + }, + } { + seedsStrings := make([][2]string, len(s)) + for i, h := range s { + seedsStrings[i] = [2]string{h.Name, h.Value} + } + data, err := json.Marshal(seedsStrings) + require.NoError(f, err) + corpus.Add(data) + } + + f.Fuzz(func(t *testing.T, data []byte) { + // Header fields are encoded as JSON (a [][2]string of [name, value] pairs) rather than as + // QPACK-encoded bytes. This bypasses the QPACK decoder intentionally: QPACK is fuzzed + // separately (in the qpack package). + const maxPairs = 1000 + const maxHeaderBytes = 50_000 + var pairs [][2]string + if err := json.Unmarshal(data, &pairs); err != nil { + return + } + if len(pairs) > maxPairs { + // don't fuzz too many header fields all at once + return + } + headers := make([]qpack.HeaderField, len(pairs)) + for i, p := range pairs { + headers[i] = qpack.HeaderField{Name: p[0], Value: p[1]} + } + + if req, err := requestFromHeaders(decodeFromSlice(headers), maxHeaderBytes, nil); err == nil { + require.NotEmpty(t, req.Method, "request has empty Method") + require.NotNil(t, req.URL, "request has nil URL") + require.NotEmpty(t, req.Proto, "request has empty Proto") + require.Truef(t, req.ProtoMajor == 3 && req.ProtoMinor == 0, "expected HTTP/3.0, got %d.%d", req.ProtoMajor, req.ProtoMinor) + require.GreaterOrEqualf(t, req.ContentLength, int64(-1), "invalid ContentLength: %d", req.ContentLength) + require.NotNil(t, req.Header, "request has nil Header map") + if req.Method == http.MethodConnect && req.Proto == "HTTP/3.0" { + // regular CONNECT: :path must be empty, :authority must be set + require.Empty(t, req.URL.Path, "CONNECT request has non-empty URL.Path") + } + if req.Method != http.MethodConnect { + require.NotEmpty(t, req.Host, "non-CONNECT request has empty Host") + require.NotEmpty(t, req.RequestURI, "non-CONNECT request has empty RequestURI") + } + requireValidFuzzHeader(t, req.Header, "request") + } + + rsp := &http.Response{} + if err := updateResponseFromHeaders(rsp, decodeFromSlice(headers), maxHeaderBytes, nil); err == nil { + require.Equalf(t, "HTTP/3.0", rsp.Proto, "expected Proto HTTP/3.0, got %q", rsp.Proto) + require.Equalf(t, 3, rsp.ProtoMajor, "expected ProtoMajor 3, got %d", rsp.ProtoMajor) + require.GreaterOrEqualf(t, rsp.ContentLength, int64(-1), "invalid ContentLength: %d", rsp.ContentLength) + require.NotNil(t, rsp.Header, "response has nil Header map") + require.NotEmpty(t, rsp.Status, "response has empty Status") + requireValidFuzzHeader(t, rsp.Header, "response") + } + + if trailers, err := parseTrailers(decodeFromSlice(headers), maxHeaderBytes, nil); err == nil { + for name := range trailers { + require.Falsef(t, len(name) > 0 && name[0] == ':', "trailer contains pseudo header %q", name) + } + requireValidFuzzTrailer(t, trailers) + } + }) +} + +func requireValidFuzzHeader(t *testing.T, h http.Header, context string) { + t.Helper() + for name, values := range h { + require.Truef(t, httpguts.ValidHeaderFieldName(name), "%s contains invalid header field name %q", context, name) + for _, value := range values { + require.Truef(t, httpguts.ValidHeaderFieldValue(value), "%s contains invalid header field value for %q: %q", context, name, value) + } + } + for _, name := range invalidHeaderFields { + require.Emptyf(t, h.Get(name), "%s contains connection-specific header %q", context, name) + } + if te := h.Values("Te"); len(te) > 0 { + for _, value := range te { + require.Equalf(t, "trailers", value, "%s contains invalid TE header field value: %q", context, value) + } + } +} + +func requireValidFuzzTrailer(t *testing.T, h http.Header) { + t.Helper() + requireValidFuzzHeader(t, h, "trailer") + for name := range h { + require.Truef(t, httpguts.ValidTrailerHeader(name), "trailer contains invalid trailer field name %q", name) + } +} diff --git a/third_party/quic-go/http3/http3_helper_test.go b/third_party/quic-go/http3/http3_helper_test.go new file mode 100644 index 0000000..ee92055 --- /dev/null +++ b/third_party/quic-go/http3/http3_helper_test.go @@ -0,0 +1,353 @@ +package http3 + +import ( + "bytes" + "context" + "crypto" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "errors" + "io" + "math/big" + "net" + "net/http" + "os" + "reflect" + "strconv" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + "github.com/quic-go/qpack" + + "github.com/stretchr/testify/require" +) + +// maxByteCount is the maximum value of a ByteCount +const maxByteCount = uint64(1<<62 - 1) + +func newUDPConnLocalhost(t testing.TB) *net.UDPConn { + t.Helper() + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + t.Cleanup(func() { conn.Close() }) + return conn +} + +func scaleDuration(t time.Duration) time.Duration { + scaleFactor := 1 + if f, err := strconv.Atoi(os.Getenv("TIMESCALE_FACTOR")); err == nil { // parsing "" errors, so this works fine if the env is not set + scaleFactor = f + } + if scaleFactor == 0 { + panic("TIMESCALE_FACTOR is 0") + } + return time.Duration(scaleFactor) * t +} + +var tlsConfig, tlsClientConfig *tls.Config + +func init() { + ca, caPrivateKey, err := generateCA() + if err != nil { + panic(err) + } + leafCert, leafPrivateKey, err := generateLeafCert(ca, caPrivateKey) + if err != nil { + panic(err) + } + tlsConfig = &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: [][]byte{leafCert.Raw}, + PrivateKey: leafPrivateKey, + }}, + NextProtos: []string{NextProtoH3}, + } + + root := x509.NewCertPool() + root.AddCert(ca) + tlsClientConfig = &tls.Config{ + ServerName: "localhost", + RootCAs: root, + NextProtos: []string{NextProtoH3}, + } +} + +func generateCA() (*x509.Certificate, crypto.PrivateKey, error) { + certTempl := &x509.Certificate{ + SerialNumber: big.NewInt(2019), + Subject: pkix.Name{}, + NotBefore: time.Now(), + NotAfter: time.Now().Add(24 * time.Hour), + IsCA: true, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + BasicConstraintsValid: true, + } + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + caBytes, err := x509.CreateCertificate(rand.Reader, certTempl, certTempl, pub, priv) + if err != nil { + return nil, nil, err + } + ca, err := x509.ParseCertificate(caBytes) + if err != nil { + return nil, nil, err + } + return ca, priv, nil +} + +func generateLeafCert(ca *x509.Certificate, caPriv crypto.PrivateKey) (*x509.Certificate, crypto.PrivateKey, error) { + certTempl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + DNSNames: []string{"localhost"}, + IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)}, + NotBefore: time.Now(), + NotAfter: time.Now().Add(24 * time.Hour), + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, + KeyUsage: x509.KeyUsageDigitalSignature, + } + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + certBytes, err := x509.CreateCertificate(rand.Reader, certTempl, ca, pub, caPriv) + if err != nil { + return nil, nil, err + } + cert, err := x509.ParseCertificate(certBytes) + if err != nil { + return nil, nil, err + } + return cert, priv, nil +} + +func getTLSConfig() *tls.Config { return tlsConfig.Clone() } +func getTLSClientConfig() *tls.Config { return tlsClientConfig.Clone() } + +type qlogTrace struct { + recorder qlogwriter.Recorder +} + +func (t *qlogTrace) SupportsSchemas(schema string) bool { return true } + +func (t *qlogTrace) AddProducer() qlogwriter.Recorder { + return t.recorder +} + +type connPairOpts struct { + clientRecorder qlogwriter.Recorder + serverRecorder qlogwriter.Recorder + serverBidiStreamLimit int64 + enableDatagrams bool +} + +type connPairOpt func(*connPairOpts) + +func withClientRecorder(r qlogwriter.Recorder) connPairOpt { + return func(o *connPairOpts) { o.clientRecorder = r } +} + +func withServerRecorder(r qlogwriter.Recorder) connPairOpt { + return func(o *connPairOpts) { o.serverRecorder = r } +} + +func withDatagrams() connPairOpt { + return func(o *connPairOpts) { o.enableDatagrams = true } +} + +func withServerBidiStreamLimit(limit int64) connPairOpt { + return func(o *connPairOpts) { o.serverBidiStreamLimit = limit } +} + +func newConnPair(t *testing.T, opts ...connPairOpt) (client, server *quic.Conn) { + t.Helper() + + var o connPairOpts + for _, opt := range opts { + opt(&o) + } + + ln, err := quic.ListenEarly( + newUDPConnLocalhost(t), + getTLSConfig(), + &quic.Config{ + InitialStreamReceiveWindow: maxByteCount, + InitialConnectionReceiveWindow: maxByteCount, + MaxIncomingStreams: o.serverBidiStreamLimit, + EnableDatagrams: o.enableDatagrams, + Tracer: func(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + return &qlogTrace{recorder: o.serverRecorder} + }, + }, + ) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + cl, err := quic.DialEarly( + ctx, + newUDPConnLocalhost(t), + ln.Addr(), + getTLSClientConfig(), + &quic.Config{ + EnableDatagrams: o.enableDatagrams, + Tracer: func(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + return &qlogTrace{recorder: o.clientRecorder} + }, + }, + ) + require.NoError(t, err) + t.Cleanup(func() { cl.CloseWithError(0, "") }) + + conn, err := ln.Accept(ctx) + require.NoError(t, err) + t.Cleanup(func() { conn.CloseWithError(0, "") }) + select { + case <-conn.HandshakeComplete(): + case <-ctx.Done(): + t.Fatal("timeout") + } + return cl, conn +} + +type quicReceiveStream interface { + io.Reader + SetReadDeadline(time.Time) error +} + +func expectStreamReadReset(t *testing.T, str quicReceiveStream, errCode quic.StreamErrorCode) { + t.Helper() + + str.SetReadDeadline(time.Now().Add(time.Second)) + _, err := str.Read([]byte{0}) + require.Error(t, err) + if errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatal("didn't receive a stream reset") + } + var strErr *quic.StreamError + require.ErrorAs(t, err, &strErr) + require.Equal(t, errCode, strErr.ErrorCode) +} + +type quicSendStream interface { + io.Writer + Context() context.Context +} + +func expectStreamWriteReset(t *testing.T, str quicSendStream, errCode quic.StreamErrorCode) { + t.Helper() + + select { + case <-str.Context().Done(): + case <-time.After(time.Second): + t.Fatal("timeout") + } + _, err := str.Write([]byte{0}) + require.Error(t, err) + var strErr *quic.StreamError + require.ErrorAs(t, err, &strErr) + require.Equal(t, errCode, strErr.ErrorCode) +} + +func encodeRequest(t *testing.T, req *http.Request) []byte { + t.Helper() + + var buf bytes.Buffer + rw := newRequestWriter() + require.NoError(t, rw.WriteRequestHeader(&buf, req, false, 0, nil)) + if req.Body != nil { + body, err := io.ReadAll(req.Body) + require.NoError(t, err) + buf.Write((&dataFrame{Length: uint64(len(body))}).Append(nil)) + buf.Write(body) + } + return buf.Bytes() +} + +func decodeHeader(t *testing.T, r io.Reader) map[string][]string { + t.Helper() + + fields := make(map[string][]string) + frame, err := (&frameParser{r: r}).ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &headersFrame{}, frame) + headersFrame := frame.(*headersFrame) + data := make([]byte, headersFrame.Length) + _, err = io.ReadFull(r, data) + require.NoError(t, err) + hfs := decodeQpackHeaderFields(t, data) + for _, p := range hfs { + fields[p.Name] = append(fields[p.Name], p.Value) + } + return fields +} + +func decodeQpackHeaderFields(t *testing.T, data []byte) []qpack.HeaderField { + t.Helper() + + decoder := qpack.NewDecoder() + decodeFn := decoder.Decode(data) + var hfs []qpack.HeaderField + for { + hf, err := decodeFn() + if err == io.EOF { + break + } + require.NoError(t, err) + hfs = append(hfs, hf) + } + return hfs +} + +// filterQlogEventsForFrame filters the events for the given frame type, +// for both FrameCreated and FrameParsed events. +// It returns the events that match the given frame type. +func filterQlogEventsForFrame(events []qlogwriter.Event, frame any) []qlogwriter.Event { + var filtered []qlogwriter.Event + for _, ev := range events { + switch e := ev.(type) { + case qlog.FrameCreated: + if reflect.TypeOf(e.Frame.Frame) == reflect.TypeOf(frame) { + filtered = append(filtered, ev) + } + case qlog.FrameParsed: + if reflect.TypeOf(e.Frame.Frame) == reflect.TypeOf(frame) { + filtered = append(filtered, ev) + } + } + } + return filtered +} + +func expectedFrameLength(t *testing.T, frame any) (length, payloadLength int) { + t.Helper() + + switch f := frame.(type) { + case *dataFrame: + return len(f.Append(nil)) + int(f.Length), int(f.Length) + case *headersFrame: + return len(f.Append(nil)) + int(f.Length), int(f.Length) + case *goAwayFrame: + return len(f.Append(nil)), quicvarint.Len(uint64(f.StreamID)) + case *settingsFrame: + data := f.Append(nil) + r := bytes.NewReader(data) + _, err := quicvarint.Read(r) // type + require.NoError(t, err) + _, err = quicvarint.Read(r) // length + require.NoError(t, err) + return len(data), r.Len() + default: + t.Fatalf("unexpected frame type: %T", frame) + } + panic("unreachable") +} diff --git a/third_party/quic-go/http3/internal/testdata/cert.go b/third_party/quic-go/http3/internal/testdata/cert.go new file mode 100644 index 0000000..a05a576 --- /dev/null +++ b/third_party/quic-go/http3/internal/testdata/cert.go @@ -0,0 +1,28 @@ +package testdata + +import ( + "crypto/tls" + "crypto/x509" + + quictestdata "github.com/apernet/quic-go/internal/testdata" +) + +// GetCertificatePaths returns the paths to certificate and key +func GetCertificatePaths() (string, string) { + return quictestdata.GetCertificatePaths() +} + +// GetTLSConfig returns a TLS config for localhost. +func GetTLSConfig() *tls.Config { + return quictestdata.GetTLSConfig() +} + +// AddRootCA adds the root CA certificate to a cert pool +func AddRootCA(certPool *x509.CertPool) { + quictestdata.AddRootCA(certPool) +} + +// GetRootCA returns an x509.CertPool containing (only) the CA certificate +func GetRootCA() *x509.CertPool { + return quictestdata.GetRootCA() +} diff --git a/third_party/quic-go/http3/internal/testdata/cert_test.go b/third_party/quic-go/http3/internal/testdata/cert_test.go new file mode 100644 index 0000000..e3e4a79 --- /dev/null +++ b/third_party/quic-go/http3/internal/testdata/cert_test.go @@ -0,0 +1,31 @@ +package testdata + +import ( + "crypto/tls" + "io" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCertificates(t *testing.T) { + ln, err := tls.Listen("tcp", "localhost:0", GetTLSConfig()) + require.NoError(t, err) + + go func() { + conn, err := ln.Accept() + require.NoError(t, err) + defer conn.Close() + _, err = conn.Write([]byte("foobar")) + require.NoError(t, err) + }() + + conn, err := tls.Dial("tcp", ln.Addr().String(), &tls.Config{ + RootCAs: GetRootCA(), + ServerName: "localhost", + }) + require.NoError(t, err) + data, err := io.ReadAll(conn) + require.NoError(t, err) + require.Equal(t, "foobar", string(data)) +} diff --git a/third_party/quic-go/http3/ip_addr.go b/third_party/quic-go/http3/ip_addr.go new file mode 100644 index 0000000..876a1e3 --- /dev/null +++ b/third_party/quic-go/http3/ip_addr.go @@ -0,0 +1,48 @@ +package http3 + +import ( + "net" + "strings" +) + +// An addrList represents a list of network endpoint addresses. +// Copy from [net.addrList] and change type from [net.Addr] to [net.IPAddr] +type addrList []net.IPAddr + +// isIPv4 reports whether addr contains an IPv4 address. +func isIPv4(addr net.IPAddr) bool { + return addr.IP.To4() != nil +} + +// isNotIPv4 reports whether addr does not contain an IPv4 address. +func isNotIPv4(addr net.IPAddr) bool { return !isIPv4(addr) } + +// forResolve returns the most appropriate address in address for +// a call to ResolveTCPAddr, ResolveUDPAddr, or ResolveIPAddr. +// IPv4 is preferred, unless addr contains an IPv6 literal. +func (addrs addrList) forResolve(network, addr string) net.IPAddr { + var want6 bool + switch network { + case "ip": + // IPv6 literal (addr does NOT contain a port) + want6 = strings.ContainsRune(addr, ':') + case "tcp", "udp": + // IPv6 literal. (addr contains a port, so look for '[') + want6 = strings.ContainsRune(addr, '[') + } + if want6 { + return addrs.first(isNotIPv4) + } + return addrs.first(isIPv4) +} + +// first returns the first address which satisfies strategy, or if +// none do, then the first address of any kind. +func (addrs addrList) first(strategy func(net.IPAddr) bool) net.IPAddr { + for _, addr := range addrs { + if strategy(addr) { + return addr + } + } + return addrs[0] +} diff --git a/third_party/quic-go/http3/mock_clientconn_test.go b/third_party/quic-go/http3/mock_clientconn_test.go new file mode 100644 index 0000000..a2dee22 --- /dev/null +++ b/third_party/quic-go/http3/mock_clientconn_test.go @@ -0,0 +1,157 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/http3 (interfaces: TestClientConnInterface) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -mock_names=TestClientConnInterface=MockClientConn -package http3 -destination mock_clientconn_test.go github.com/apernet/quic-go/http3 TestClientConnInterface +// + +// Package http3 is a generated GoMock package. +package http3 + +import ( + context "context" + http "net/http" + reflect "reflect" + + quic "github.com/apernet/quic-go" + gomock "go.uber.org/mock/gomock" +) + +// MockClientConn is a mock of TestClientConnInterface interface. +type MockClientConn struct { + ctrl *gomock.Controller + recorder *MockClientConnMockRecorder + isgomock struct{} +} + +// MockClientConnMockRecorder is the mock recorder for MockClientConn. +type MockClientConnMockRecorder struct { + mock *MockClientConn +} + +// NewMockClientConn creates a new mock instance. +func NewMockClientConn(ctrl *gomock.Controller) *MockClientConn { + mock := &MockClientConn{ctrl: ctrl} + mock.recorder = &MockClientConnMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockClientConn) EXPECT() *MockClientConnMockRecorder { + return m.recorder +} + +// OpenRequestStream mocks base method. +func (m *MockClientConn) OpenRequestStream(arg0 context.Context) (*RequestStream, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "OpenRequestStream", arg0) + ret0, _ := ret[0].(*RequestStream) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// OpenRequestStream indicates an expected call of OpenRequestStream. +func (mr *MockClientConnMockRecorder) OpenRequestStream(arg0 any) *MockClientConnOpenRequestStreamCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OpenRequestStream", reflect.TypeOf((*MockClientConn)(nil).OpenRequestStream), arg0) + return &MockClientConnOpenRequestStreamCall{Call: call} +} + +// MockClientConnOpenRequestStreamCall wrap *gomock.Call +type MockClientConnOpenRequestStreamCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockClientConnOpenRequestStreamCall) Return(arg0 *RequestStream, arg1 error) *MockClientConnOpenRequestStreamCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockClientConnOpenRequestStreamCall) Do(f func(context.Context) (*RequestStream, error)) *MockClientConnOpenRequestStreamCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockClientConnOpenRequestStreamCall) DoAndReturn(f func(context.Context) (*RequestStream, error)) *MockClientConnOpenRequestStreamCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// RoundTrip mocks base method. +func (m *MockClientConn) RoundTrip(arg0 *http.Request) (*http.Response, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "RoundTrip", arg0) + ret0, _ := ret[0].(*http.Response) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// RoundTrip indicates an expected call of RoundTrip. +func (mr *MockClientConnMockRecorder) RoundTrip(arg0 any) *MockClientConnRoundTripCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RoundTrip", reflect.TypeOf((*MockClientConn)(nil).RoundTrip), arg0) + return &MockClientConnRoundTripCall{Call: call} +} + +// MockClientConnRoundTripCall wrap *gomock.Call +type MockClientConnRoundTripCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockClientConnRoundTripCall) Return(arg0 *http.Response, arg1 error) *MockClientConnRoundTripCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockClientConnRoundTripCall) Do(f func(*http.Request) (*http.Response, error)) *MockClientConnRoundTripCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockClientConnRoundTripCall) DoAndReturn(f func(*http.Request) (*http.Response, error)) *MockClientConnRoundTripCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// handleUnidirectionalStream mocks base method. +func (m *MockClientConn) handleUnidirectionalStream(arg0 *quic.ReceiveStream) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "handleUnidirectionalStream", arg0) +} + +// handleUnidirectionalStream indicates an expected call of handleUnidirectionalStream. +func (mr *MockClientConnMockRecorder) handleUnidirectionalStream(arg0 any) *MockClientConnhandleUnidirectionalStreamCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "handleUnidirectionalStream", reflect.TypeOf((*MockClientConn)(nil).handleUnidirectionalStream), arg0) + return &MockClientConnhandleUnidirectionalStreamCall{Call: call} +} + +// MockClientConnhandleUnidirectionalStreamCall wrap *gomock.Call +type MockClientConnhandleUnidirectionalStreamCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockClientConnhandleUnidirectionalStreamCall) Return() *MockClientConnhandleUnidirectionalStreamCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockClientConnhandleUnidirectionalStreamCall) Do(f func(*quic.ReceiveStream)) *MockClientConnhandleUnidirectionalStreamCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockClientConnhandleUnidirectionalStreamCall) DoAndReturn(f func(*quic.ReceiveStream)) *MockClientConnhandleUnidirectionalStreamCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/http3/mock_datagram_stream_test.go b/third_party/quic-go/http3/mock_datagram_stream_test.go new file mode 100644 index 0000000..7e445ac --- /dev/null +++ b/third_party/quic-go/http3/mock_datagram_stream_test.go @@ -0,0 +1,536 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/http3 (interfaces: DatagramStream) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -mock_names=DatagramStream=MockDatagramStream -package http3 -destination mock_datagram_stream_test.go github.com/apernet/quic-go/http3 DatagramStream +// + +// Package http3 is a generated GoMock package. +package http3 + +import ( + context "context" + reflect "reflect" + time "time" + + quic "github.com/apernet/quic-go" + gomock "go.uber.org/mock/gomock" +) + +// MockDatagramStream is a mock of DatagramStream interface. +type MockDatagramStream struct { + ctrl *gomock.Controller + recorder *MockDatagramStreamMockRecorder + isgomock struct{} +} + +// MockDatagramStreamMockRecorder is the mock recorder for MockDatagramStream. +type MockDatagramStreamMockRecorder struct { + mock *MockDatagramStream +} + +// NewMockDatagramStream creates a new mock instance. +func NewMockDatagramStream(ctrl *gomock.Controller) *MockDatagramStream { + mock := &MockDatagramStream{ctrl: ctrl} + mock.recorder = &MockDatagramStreamMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockDatagramStream) EXPECT() *MockDatagramStreamMockRecorder { + return m.recorder +} + +// CancelRead mocks base method. +func (m *MockDatagramStream) CancelRead(arg0 quic.StreamErrorCode) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "CancelRead", arg0) +} + +// CancelRead indicates an expected call of CancelRead. +func (mr *MockDatagramStreamMockRecorder) CancelRead(arg0 any) *MockDatagramStreamCancelReadCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CancelRead", reflect.TypeOf((*MockDatagramStream)(nil).CancelRead), arg0) + return &MockDatagramStreamCancelReadCall{Call: call} +} + +// MockDatagramStreamCancelReadCall wrap *gomock.Call +type MockDatagramStreamCancelReadCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamCancelReadCall) Return() *MockDatagramStreamCancelReadCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamCancelReadCall) Do(f func(quic.StreamErrorCode)) *MockDatagramStreamCancelReadCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamCancelReadCall) DoAndReturn(f func(quic.StreamErrorCode)) *MockDatagramStreamCancelReadCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// CancelWrite mocks base method. +func (m *MockDatagramStream) CancelWrite(arg0 quic.StreamErrorCode) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "CancelWrite", arg0) +} + +// CancelWrite indicates an expected call of CancelWrite. +func (mr *MockDatagramStreamMockRecorder) CancelWrite(arg0 any) *MockDatagramStreamCancelWriteCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CancelWrite", reflect.TypeOf((*MockDatagramStream)(nil).CancelWrite), arg0) + return &MockDatagramStreamCancelWriteCall{Call: call} +} + +// MockDatagramStreamCancelWriteCall wrap *gomock.Call +type MockDatagramStreamCancelWriteCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamCancelWriteCall) Return() *MockDatagramStreamCancelWriteCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamCancelWriteCall) Do(f func(quic.StreamErrorCode)) *MockDatagramStreamCancelWriteCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamCancelWriteCall) DoAndReturn(f func(quic.StreamErrorCode)) *MockDatagramStreamCancelWriteCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Close mocks base method. +func (m *MockDatagramStream) Close() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Close") + ret0, _ := ret[0].(error) + return ret0 +} + +// Close indicates an expected call of Close. +func (mr *MockDatagramStreamMockRecorder) Close() *MockDatagramStreamCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockDatagramStream)(nil).Close)) + return &MockDatagramStreamCloseCall{Call: call} +} + +// MockDatagramStreamCloseCall wrap *gomock.Call +type MockDatagramStreamCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamCloseCall) Return(arg0 error) *MockDatagramStreamCloseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamCloseCall) Do(f func() error) *MockDatagramStreamCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamCloseCall) DoAndReturn(f func() error) *MockDatagramStreamCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Context mocks base method. +func (m *MockDatagramStream) Context() context.Context { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Context") + ret0, _ := ret[0].(context.Context) + return ret0 +} + +// Context indicates an expected call of Context. +func (mr *MockDatagramStreamMockRecorder) Context() *MockDatagramStreamContextCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Context", reflect.TypeOf((*MockDatagramStream)(nil).Context)) + return &MockDatagramStreamContextCall{Call: call} +} + +// MockDatagramStreamContextCall wrap *gomock.Call +type MockDatagramStreamContextCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamContextCall) Return(arg0 context.Context) *MockDatagramStreamContextCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamContextCall) Do(f func() context.Context) *MockDatagramStreamContextCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamContextCall) DoAndReturn(f func() context.Context) *MockDatagramStreamContextCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// QUICStream mocks base method. +func (m *MockDatagramStream) QUICStream() *quic.Stream { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "QUICStream") + ret0, _ := ret[0].(*quic.Stream) + return ret0 +} + +// QUICStream indicates an expected call of QUICStream. +func (mr *MockDatagramStreamMockRecorder) QUICStream() *MockDatagramStreamQUICStreamCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "QUICStream", reflect.TypeOf((*MockDatagramStream)(nil).QUICStream)) + return &MockDatagramStreamQUICStreamCall{Call: call} +} + +// MockDatagramStreamQUICStreamCall wrap *gomock.Call +type MockDatagramStreamQUICStreamCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamQUICStreamCall) Return(arg0 *quic.Stream) *MockDatagramStreamQUICStreamCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamQUICStreamCall) Do(f func() *quic.Stream) *MockDatagramStreamQUICStreamCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamQUICStreamCall) DoAndReturn(f func() *quic.Stream) *MockDatagramStreamQUICStreamCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Read mocks base method. +func (m *MockDatagramStream) Read(p []byte) (int, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Read", p) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Read indicates an expected call of Read. +func (mr *MockDatagramStreamMockRecorder) Read(p any) *MockDatagramStreamReadCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Read", reflect.TypeOf((*MockDatagramStream)(nil).Read), p) + return &MockDatagramStreamReadCall{Call: call} +} + +// MockDatagramStreamReadCall wrap *gomock.Call +type MockDatagramStreamReadCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamReadCall) Return(n int, err error) *MockDatagramStreamReadCall { + c.Call = c.Call.Return(n, err) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamReadCall) Do(f func([]byte) (int, error)) *MockDatagramStreamReadCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamReadCall) DoAndReturn(f func([]byte) (int, error)) *MockDatagramStreamReadCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReceiveDatagram mocks base method. +func (m *MockDatagramStream) ReceiveDatagram(ctx context.Context) ([]byte, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReceiveDatagram", ctx) + ret0, _ := ret[0].([]byte) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ReceiveDatagram indicates an expected call of ReceiveDatagram. +func (mr *MockDatagramStreamMockRecorder) ReceiveDatagram(ctx any) *MockDatagramStreamReceiveDatagramCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReceiveDatagram", reflect.TypeOf((*MockDatagramStream)(nil).ReceiveDatagram), ctx) + return &MockDatagramStreamReceiveDatagramCall{Call: call} +} + +// MockDatagramStreamReceiveDatagramCall wrap *gomock.Call +type MockDatagramStreamReceiveDatagramCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamReceiveDatagramCall) Return(arg0 []byte, arg1 error) *MockDatagramStreamReceiveDatagramCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamReceiveDatagramCall) Do(f func(context.Context) ([]byte, error)) *MockDatagramStreamReceiveDatagramCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamReceiveDatagramCall) DoAndReturn(f func(context.Context) ([]byte, error)) *MockDatagramStreamReceiveDatagramCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SendDatagram mocks base method. +func (m *MockDatagramStream) SendDatagram(b []byte) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SendDatagram", b) + ret0, _ := ret[0].(error) + return ret0 +} + +// SendDatagram indicates an expected call of SendDatagram. +func (mr *MockDatagramStreamMockRecorder) SendDatagram(b any) *MockDatagramStreamSendDatagramCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SendDatagram", reflect.TypeOf((*MockDatagramStream)(nil).SendDatagram), b) + return &MockDatagramStreamSendDatagramCall{Call: call} +} + +// MockDatagramStreamSendDatagramCall wrap *gomock.Call +type MockDatagramStreamSendDatagramCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamSendDatagramCall) Return(arg0 error) *MockDatagramStreamSendDatagramCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamSendDatagramCall) Do(f func([]byte) error) *MockDatagramStreamSendDatagramCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamSendDatagramCall) DoAndReturn(f func([]byte) error) *MockDatagramStreamSendDatagramCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetDeadline mocks base method. +func (m *MockDatagramStream) SetDeadline(arg0 time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetDeadline", arg0) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetDeadline indicates an expected call of SetDeadline. +func (mr *MockDatagramStreamMockRecorder) SetDeadline(arg0 any) *MockDatagramStreamSetDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetDeadline", reflect.TypeOf((*MockDatagramStream)(nil).SetDeadline), arg0) + return &MockDatagramStreamSetDeadlineCall{Call: call} +} + +// MockDatagramStreamSetDeadlineCall wrap *gomock.Call +type MockDatagramStreamSetDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamSetDeadlineCall) Return(arg0 error) *MockDatagramStreamSetDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamSetDeadlineCall) Do(f func(time.Time) error) *MockDatagramStreamSetDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamSetDeadlineCall) DoAndReturn(f func(time.Time) error) *MockDatagramStreamSetDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetReadDeadline mocks base method. +func (m *MockDatagramStream) SetReadDeadline(arg0 time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetReadDeadline", arg0) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetReadDeadline indicates an expected call of SetReadDeadline. +func (mr *MockDatagramStreamMockRecorder) SetReadDeadline(arg0 any) *MockDatagramStreamSetReadDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetReadDeadline", reflect.TypeOf((*MockDatagramStream)(nil).SetReadDeadline), arg0) + return &MockDatagramStreamSetReadDeadlineCall{Call: call} +} + +// MockDatagramStreamSetReadDeadlineCall wrap *gomock.Call +type MockDatagramStreamSetReadDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamSetReadDeadlineCall) Return(arg0 error) *MockDatagramStreamSetReadDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamSetReadDeadlineCall) Do(f func(time.Time) error) *MockDatagramStreamSetReadDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamSetReadDeadlineCall) DoAndReturn(f func(time.Time) error) *MockDatagramStreamSetReadDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetWriteDeadline mocks base method. +func (m *MockDatagramStream) SetWriteDeadline(arg0 time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetWriteDeadline", arg0) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetWriteDeadline indicates an expected call of SetWriteDeadline. +func (mr *MockDatagramStreamMockRecorder) SetWriteDeadline(arg0 any) *MockDatagramStreamSetWriteDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetWriteDeadline", reflect.TypeOf((*MockDatagramStream)(nil).SetWriteDeadline), arg0) + return &MockDatagramStreamSetWriteDeadlineCall{Call: call} +} + +// MockDatagramStreamSetWriteDeadlineCall wrap *gomock.Call +type MockDatagramStreamSetWriteDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamSetWriteDeadlineCall) Return(arg0 error) *MockDatagramStreamSetWriteDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamSetWriteDeadlineCall) Do(f func(time.Time) error) *MockDatagramStreamSetWriteDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamSetWriteDeadlineCall) DoAndReturn(f func(time.Time) error) *MockDatagramStreamSetWriteDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// StreamID mocks base method. +func (m *MockDatagramStream) StreamID() quic.StreamID { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "StreamID") + ret0, _ := ret[0].(quic.StreamID) + return ret0 +} + +// StreamID indicates an expected call of StreamID. +func (mr *MockDatagramStreamMockRecorder) StreamID() *MockDatagramStreamStreamIDCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StreamID", reflect.TypeOf((*MockDatagramStream)(nil).StreamID)) + return &MockDatagramStreamStreamIDCall{Call: call} +} + +// MockDatagramStreamStreamIDCall wrap *gomock.Call +type MockDatagramStreamStreamIDCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamStreamIDCall) Return(arg0 quic.StreamID) *MockDatagramStreamStreamIDCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamStreamIDCall) Do(f func() quic.StreamID) *MockDatagramStreamStreamIDCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamStreamIDCall) DoAndReturn(f func() quic.StreamID) *MockDatagramStreamStreamIDCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Write mocks base method. +func (m *MockDatagramStream) Write(p []byte) (int, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Write", p) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Write indicates an expected call of Write. +func (mr *MockDatagramStreamMockRecorder) Write(p any) *MockDatagramStreamWriteCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Write", reflect.TypeOf((*MockDatagramStream)(nil).Write), p) + return &MockDatagramStreamWriteCall{Call: call} +} + +// MockDatagramStreamWriteCall wrap *gomock.Call +type MockDatagramStreamWriteCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockDatagramStreamWriteCall) Return(n int, err error) *MockDatagramStreamWriteCall { + c.Call = c.Call.Return(n, err) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockDatagramStreamWriteCall) Do(f func([]byte) (int, error)) *MockDatagramStreamWriteCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockDatagramStreamWriteCall) DoAndReturn(f func([]byte) (int, error)) *MockDatagramStreamWriteCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/http3/mock_quic_listener_test.go b/third_party/quic-go/http3/mock_quic_listener_test.go new file mode 100644 index 0000000..420dc32 --- /dev/null +++ b/third_party/quic-go/http3/mock_quic_listener_test.go @@ -0,0 +1,158 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/http3 (interfaces: QUICListener) +// +// Generated by this command: +// +// mockgen -typed -package http3 -destination mock_quic_listener_test.go github.com/apernet/quic-go/http3 QUICListener +// + +// Package http3 is a generated GoMock package. +package http3 + +import ( + context "context" + net "net" + reflect "reflect" + + quic "github.com/apernet/quic-go" + gomock "go.uber.org/mock/gomock" +) + +// MockQUICListener is a mock of QUICListener interface. +type MockQUICListener struct { + ctrl *gomock.Controller + recorder *MockQUICListenerMockRecorder + isgomock struct{} +} + +// MockQUICListenerMockRecorder is the mock recorder for MockQUICListener. +type MockQUICListenerMockRecorder struct { + mock *MockQUICListener +} + +// NewMockQUICListener creates a new mock instance. +func NewMockQUICListener(ctrl *gomock.Controller) *MockQUICListener { + mock := &MockQUICListener{ctrl: ctrl} + mock.recorder = &MockQUICListenerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockQUICListener) EXPECT() *MockQUICListenerMockRecorder { + return m.recorder +} + +// Accept mocks base method. +func (m *MockQUICListener) Accept(arg0 context.Context) (*quic.Conn, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Accept", arg0) + ret0, _ := ret[0].(*quic.Conn) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Accept indicates an expected call of Accept. +func (mr *MockQUICListenerMockRecorder) Accept(arg0 any) *MockQUICListenerAcceptCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Accept", reflect.TypeOf((*MockQUICListener)(nil).Accept), arg0) + return &MockQUICListenerAcceptCall{Call: call} +} + +// MockQUICListenerAcceptCall wrap *gomock.Call +type MockQUICListenerAcceptCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockQUICListenerAcceptCall) Return(arg0 *quic.Conn, arg1 error) *MockQUICListenerAcceptCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockQUICListenerAcceptCall) Do(f func(context.Context) (*quic.Conn, error)) *MockQUICListenerAcceptCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockQUICListenerAcceptCall) DoAndReturn(f func(context.Context) (*quic.Conn, error)) *MockQUICListenerAcceptCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Addr mocks base method. +func (m *MockQUICListener) Addr() net.Addr { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Addr") + ret0, _ := ret[0].(net.Addr) + return ret0 +} + +// Addr indicates an expected call of Addr. +func (mr *MockQUICListenerMockRecorder) Addr() *MockQUICListenerAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Addr", reflect.TypeOf((*MockQUICListener)(nil).Addr)) + return &MockQUICListenerAddrCall{Call: call} +} + +// MockQUICListenerAddrCall wrap *gomock.Call +type MockQUICListenerAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockQUICListenerAddrCall) Return(arg0 net.Addr) *MockQUICListenerAddrCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockQUICListenerAddrCall) Do(f func() net.Addr) *MockQUICListenerAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockQUICListenerAddrCall) DoAndReturn(f func() net.Addr) *MockQUICListenerAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Close mocks base method. +func (m *MockQUICListener) Close() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Close") + ret0, _ := ret[0].(error) + return ret0 +} + +// Close indicates an expected call of Close. +func (mr *MockQUICListenerMockRecorder) Close() *MockQUICListenerCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockQUICListener)(nil).Close)) + return &MockQUICListenerCloseCall{Call: call} +} + +// MockQUICListenerCloseCall wrap *gomock.Call +type MockQUICListenerCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockQUICListenerCloseCall) Return(arg0 error) *MockQUICListenerCloseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockQUICListenerCloseCall) Do(f func() error) *MockQUICListenerCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockQUICListenerCloseCall) DoAndReturn(f func() error) *MockQUICListenerCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/http3/mockgen.go b/third_party/quic-go/http3/mockgen.go new file mode 100644 index 0000000..937dcd6 --- /dev/null +++ b/third_party/quic-go/http3/mockgen.go @@ -0,0 +1,11 @@ +//go:build gomock || generate + +package http3 + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -mock_names=TestClientConnInterface=MockClientConn -package http3 -destination mock_clientconn_test.go github.com/apernet/quic-go/http3 TestClientConnInterface" +type TestClientConnInterface = clientConn + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -mock_names=DatagramStream=MockDatagramStream -package http3 -destination mock_datagram_stream_test.go github.com/apernet/quic-go/http3 DatagramStream" +type DatagramStream = datagramStream + +//go:generate sh -c "go tool mockgen -typed -package http3 -destination mock_quic_listener_test.go github.com/apernet/quic-go/http3 QUICListener" diff --git a/third_party/quic-go/http3/qlog.go b/third_party/quic-go/http3/qlog.go new file mode 100644 index 0000000..0a31cc8 --- /dev/null +++ b/third_party/quic-go/http3/qlog.go @@ -0,0 +1,56 @@ +package http3 + +import ( + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/quic-go/qpack" +) + +func maybeQlogInvalidHeadersFrame(qlogger qlogwriter.Recorder, streamID quic.StreamID, l uint64) { + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: streamID, + Raw: qlog.RawInfo{PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.HeadersFrame{}}, + }) + } +} + +func qlogParsedHeadersFrame(qlogger qlogwriter.Recorder, streamID quic.StreamID, hf *headersFrame, hfs []qpack.HeaderField) { + headerFields := make([]qlog.HeaderField, len(hfs)) + for i, hf := range hfs { + headerFields[i] = qlog.HeaderField{ + Name: hf.Name, + Value: hf.Value, + } + } + qlogger.RecordEvent(qlog.FrameParsed{ + StreamID: streamID, + Raw: qlog.RawInfo{ + Length: int(hf.Length) + hf.headerLen, + PayloadLength: int(hf.Length), + }, + Frame: qlog.Frame{Frame: qlog.HeadersFrame{ + HeaderFields: headerFields, + }}, + }) +} + +func qlogCreatedHeadersFrame(qlogger qlogwriter.Recorder, streamID quic.StreamID, length, payloadLength int, hfs []qlog.HeaderField) { + headerFields := make([]qlog.HeaderField, len(hfs)) + for i, hf := range hfs { + headerFields[i] = qlog.HeaderField{ + Name: hf.Name, + Value: hf.Value, + } + } + qlogger.RecordEvent(qlog.FrameCreated{ + StreamID: streamID, + Raw: qlog.RawInfo{Length: length, PayloadLength: payloadLength}, + Frame: qlog.Frame{Frame: qlog.HeadersFrame{ + HeaderFields: headerFields, + }}, + }) +} diff --git a/third_party/quic-go/http3/qlog/event.go b/third_party/quic-go/http3/qlog/event.go new file mode 100644 index 0000000..e938245 --- /dev/null +++ b/third_party/quic-go/http3/qlog/event.go @@ -0,0 +1,138 @@ +package qlog + +import ( + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +type encoderHelper struct { + enc *jsontext.Encoder + err error +} + +func (h *encoderHelper) WriteToken(t jsontext.Token) { + if h.err != nil { + return + } + h.err = h.enc.WriteToken(t) +} + +type RawInfo struct { + Length int // full packet length, including header and AEAD authentication tag + PayloadLength int // length of the packet payload, excluding AEAD tag +} + +func (i RawInfo) HasValues() bool { + return i.Length != 0 || i.PayloadLength != 0 +} + +func (i RawInfo) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if i.Length != 0 { + h.WriteToken(jsontext.String("length")) + h.WriteToken(jsontext.Uint(uint64(i.Length))) + } + if i.PayloadLength != 0 { + h.WriteToken(jsontext.String("payload_length")) + h.WriteToken(jsontext.Uint(uint64(i.PayloadLength))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type FrameParsed struct { + StreamID quic.StreamID + Raw RawInfo + Frame Frame +} + +func (e FrameParsed) Name() string { return "http3:frame_parsed" } + +func (e FrameParsed) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Uint(uint64(e.StreamID))) + if e.Raw.HasValues() { + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + } + h.WriteToken(jsontext.String("frame")) + if err := e.Frame.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type FrameCreated struct { + StreamID quic.StreamID + Raw RawInfo + Frame Frame +} + +func (e FrameCreated) Name() string { return "http3:frame_created" } + +func (e FrameCreated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Uint(uint64(e.StreamID))) + if e.Raw.HasValues() { + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + } + h.WriteToken(jsontext.String("frame")) + if err := e.Frame.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type DatagramCreated struct { + QuarterStreamID uint64 + Raw RawInfo +} + +func (e DatagramCreated) Name() string { return "http3:datagram_created" } + +func (e DatagramCreated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("quarter_stream_id")) + h.WriteToken(jsontext.Uint(e.QuarterStreamID)) + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type DatagramParsed struct { + QuarterStreamID uint64 + Raw RawInfo +} + +func (e DatagramParsed) Name() string { return "http3:datagram_parsed" } + +func (e DatagramParsed) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("quarter_stream_id")) + h.WriteToken(jsontext.Uint(e.QuarterStreamID)) + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} diff --git a/third_party/quic-go/http3/qlog/event_test.go b/third_party/quic-go/http3/qlog/event_test.go new file mode 100644 index 0000000..7b4505d --- /dev/null +++ b/third_party/quic-go/http3/qlog/event_test.go @@ -0,0 +1,98 @@ +package qlog + +import ( + "bytes" + "encoding/json" + "io" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/stretchr/testify/require" +) + +type nopWriteCloserImpl struct{ io.Writer } + +func (nopWriteCloserImpl) Close() error { return nil } + +func nopWriteCloser(w io.Writer) io.WriteCloser { + return &nopWriteCloserImpl{Writer: w} +} + +func testEventEncoding(t *testing.T, ev qlogwriter.Event) (string, map[string]any) { + t.Helper() + var buf bytes.Buffer + + synctest.Test(t, func(t *testing.T) { + tr := qlogwriter.NewConnectionFileSeq( + nopWriteCloser(&buf), + true, + quic.ConnectionIDFromBytes([]byte{1, 2, 3, 4}), + []string{"http3"}, + ) + go tr.Run() + producer := tr.AddProducer() + + synctest.Wait() + time.Sleep(42 * time.Second) + + producer.RecordEvent(ev) + producer.Close() + }) + + return decode(t, buf.String()) +} + +func decode(t *testing.T, data string) (string, map[string]any) { + t.Helper() + + var result map[string]any + + lines := bytes.Split([]byte(data), []byte{'\n'}) + require.Len(t, lines, 3) // the first line is the trace header, the second line is the event, the third line is empty + require.Empty(t, lines[2]) + require.Equal(t, qlogwriter.RecordSeparator, lines[1][0], "expected record separator at start of line") + require.NoError(t, json.Unmarshal(lines[1][1:], &result)) + require.Equal(t, 42*time.Second, time.Duration(result["time"].(float64)*1e6)*time.Nanosecond) + + return result["name"].(string), result["data"].(map[string]any) +} + +func TestFrameParsedEvent(t *testing.T) { + name, ev := testEventEncoding(t, FrameParsed{ + StreamID: quic.StreamID(4), + Raw: RawInfo{ + Length: 1500, + PayloadLength: 100, + }, + Frame: Frame{Frame: &DataFrame{}}, + }) + + require.Equal(t, "http3:frame_parsed", name) + require.Equal(t, float64(4), ev["stream_id"]) + require.NotContains(t, ev, "name") + require.Contains(t, ev, "frame") +} + +func TestFrameCreatedEvent(t *testing.T) { + name, ev := testEventEncoding(t, FrameCreated{ + StreamID: quic.StreamID(8), + Raw: RawInfo{ + PayloadLength: 200, + }, + Frame: Frame{Frame: &HeadersFrame{ + HeaderFields: []HeaderField{ + {Name: ":status", Value: "200"}, + {Name: "content-type", Value: "text/html"}, + }, + }}, + }) + + require.Equal(t, "http3:frame_created", name) + require.Equal(t, float64(8), ev["stream_id"]) + require.NotContains(t, ev, "name") + require.Contains(t, ev, "frame") +} diff --git a/third_party/quic-go/http3/qlog/frame.go b/third_party/quic-go/http3/qlog/frame.go new file mode 100644 index 0000000..1a40eba --- /dev/null +++ b/third_party/quic-go/http3/qlog/frame.go @@ -0,0 +1,220 @@ +package qlog + +import ( + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +// Frame represents an HTTP/3 frame. +type Frame struct { + Frame any +} + +func (f Frame) encode(enc *jsontext.Encoder) error { + switch frame := f.Frame.(type) { + case DataFrame: + return frame.encode(enc) + case HeadersFrame: + return frame.encode(enc) + case GoAwayFrame: + return frame.encode(enc) + case SettingsFrame: + return frame.encode(enc) + case PushPromiseFrame: + return frame.encode(enc) + case CancelPushFrame: + return frame.encode(enc) + case MaxPushIDFrame: + return frame.encode(enc) + case ReservedFrame: + return frame.encode(enc) + case UnknownFrame: + return frame.encode(enc) + } + // This shouldn't happen if the code is correctly logging frames. + // Write a null token to produce valid JSON. + return enc.WriteToken(jsontext.Null) +} + +// A DataFrame is a DATA frame +type DataFrame struct{} + +func (f *DataFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("data")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type HeaderField struct { + Name string + Value string +} + +// A HeadersFrame is a HEADERS frame +type HeadersFrame struct { + HeaderFields []HeaderField +} + +func (f *HeadersFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("headers")) + if len(f.HeaderFields) > 0 { + h.WriteToken(jsontext.String("header_fields")) + h.WriteToken(jsontext.BeginArray) + for _, f := range f.HeaderFields { + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("name")) + h.WriteToken(jsontext.String(f.Name)) + h.WriteToken(jsontext.String("value")) + h.WriteToken(jsontext.String(f.Value)) + h.WriteToken(jsontext.EndObject) + } + h.WriteToken(jsontext.EndArray) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +// A GoAwayFrame is a GOAWAY frame +type GoAwayFrame struct { + StreamID quic.StreamID +} + +func (f *GoAwayFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("goaway")) + h.WriteToken(jsontext.String("id")) + h.WriteToken(jsontext.Uint(uint64(f.StreamID))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type SettingsFrame struct { + MaxFieldSectionSize int64 + Datagram *bool + ExtendedConnect *bool + Other map[uint64]uint64 +} + +func (f *SettingsFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("settings")) + h.WriteToken(jsontext.String("settings")) + h.WriteToken(jsontext.BeginArray) + if f.MaxFieldSectionSize >= 0 { + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("name")) + h.WriteToken(jsontext.String("settings_max_field_section_size")) + h.WriteToken(jsontext.String("value")) + h.WriteToken(jsontext.Uint(uint64(f.MaxFieldSectionSize))) + h.WriteToken(jsontext.EndObject) + } + if f.Datagram != nil { + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("name")) + h.WriteToken(jsontext.String("settings_h3_datagram")) + h.WriteToken(jsontext.String("value")) + h.WriteToken(jsontext.Bool(*f.Datagram)) + h.WriteToken(jsontext.EndObject) + } + if f.ExtendedConnect != nil { + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("name")) + h.WriteToken(jsontext.String("settings_enable_connect_protocol")) + h.WriteToken(jsontext.String("value")) + h.WriteToken(jsontext.Bool(*f.ExtendedConnect)) + h.WriteToken(jsontext.EndObject) + } + if len(f.Other) > 0 { + for k, v := range f.Other { + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("name")) + h.WriteToken(jsontext.String("unknown")) + h.WriteToken(jsontext.String("name_bytes")) + h.WriteToken(jsontext.Uint(k)) + h.WriteToken(jsontext.String("value")) + h.WriteToken(jsontext.Uint(v)) + h.WriteToken(jsontext.EndObject) + } + } + h.WriteToken(jsontext.EndArray) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// A PushPromiseFrame is a PUSH_PROMISE frame +type PushPromiseFrame struct{} + +func (f *PushPromiseFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("push_promise")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// A CancelPushFrame is a CANCEL_PUSH frame +type CancelPushFrame struct{} + +func (f *CancelPushFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("cancel_push")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// A MaxPushIDFrame is a MAX_PUSH_ID frame +type MaxPushIDFrame struct{} + +func (f *MaxPushIDFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("max_push_id")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// A ReservedFrame is one of the reserved frame types +type ReservedFrame struct { + Type uint64 +} + +func (f *ReservedFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("reserved")) + h.WriteToken(jsontext.String("frame_type_bytes")) + h.WriteToken(jsontext.Uint(f.Type)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// An UnknownFrame is an unknown frame type +type UnknownFrame struct { + Type uint64 +} + +func (f *UnknownFrame) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("unknown")) + h.WriteToken(jsontext.String("frame_type_bytes")) + h.WriteToken(jsontext.Uint(f.Type)) + h.WriteToken(jsontext.EndObject) + return h.err +} diff --git a/third_party/quic-go/http3/qlog/frame_test.go b/third_party/quic-go/http3/qlog/frame_test.go new file mode 100644 index 0000000..965612b --- /dev/null +++ b/third_party/quic-go/http3/qlog/frame_test.go @@ -0,0 +1,204 @@ +package qlog + +import ( + "bytes" + "encoding/json" + "testing" + + "github.com/apernet/quic-go/qlogwriter/jsontext" + + "github.com/stretchr/testify/require" +) + +func check(t *testing.T, f any, expected map[string]any) { + t.Helper() + + var buf bytes.Buffer + enc := jsontext.NewEncoder(&buf) + require.NoError(t, (Frame{Frame: f}).encode(enc)) + data := buf.Bytes() + require.True(t, json.Valid(data), "invalid JSON: %s", string(data)) + checkEncoding(t, data, expected) +} + +func checkEncoding(t *testing.T, data []byte, expected map[string]any) { + t.Helper() + + m := make(map[string]any) + require.NoError(t, json.Unmarshal(data, &m)) + require.Len(t, m, len(expected)) + + for key, value := range expected { + switch v := value.(type) { + case bool, string, map[string]any: + require.Equal(t, v, m[key]) + case int: + require.Equal(t, float64(v), m[key]) + case float64: + require.Equal(t, v, m[key]) + case []map[string]any: // used for header fields + require.Contains(t, m, key) + slice, ok := m[key].([]any) + require.True(t, ok) + require.Len(t, slice, len(v)) + for i, expectedField := range v { + field, ok := slice[i].(map[string]any) + require.True(t, ok) + require.Equal(t, expectedField, field) + } + default: + t.Fatalf("unexpected type: %T", v) + } + } +} + +func TestDataFrame(t *testing.T) { + check(t, DataFrame{}, map[string]any{ + "frame_type": "data", + }) +} + +func TestHeadersFrame(t *testing.T) { + check(t, HeadersFrame{ + HeaderFields: []HeaderField{ + {Name: ":status", Value: "200"}, + {Name: "content-type", Value: "application/json"}, + }, + }, map[string]any{ + "frame_type": "headers", + "header_fields": []map[string]any{ + {"name": ":status", "value": "200"}, + {"name": "content-type", "value": "application/json"}, + }, + }) +} + +func TestGoAwayFrame(t *testing.T) { + check(t, GoAwayFrame{StreamID: 1337}, map[string]any{ + "frame_type": "goaway", + "id": 1337, + }) +} + +func pointer[T any](v T) *T { + return &v +} + +func TestSettingsFrame(t *testing.T) { + tests := []struct { + name string + frame SettingsFrame + expected map[string]any + }{ + { + name: "datagram: true", + frame: SettingsFrame{ + MaxFieldSectionSize: -1, + Datagram: pointer(true), + }, + expected: map[string]any{ + "frame_type": "settings", + "settings": []map[string]any{{ + "name": "settings_h3_datagram", + "value": true, + }}, + }, + }, + { + name: "extended_connect: false", + frame: SettingsFrame{ + MaxFieldSectionSize: -1, + ExtendedConnect: pointer(false), + }, + expected: map[string]any{ + "frame_type": "settings", + "settings": []map[string]any{{ + "name": "settings_enable_connect_protocol", + "value": false, + }}, + }, + }, + { + name: "max_field_section_size", + frame: SettingsFrame{MaxFieldSectionSize: 1337}, + expected: map[string]any{ + "frame_type": "settings", + "settings": []map[string]any{{ + "name": "settings_max_field_section_size", + "value": float64(1337), + }}, + }, + }, + { + name: "datagram: false, extended_connect: false", + frame: SettingsFrame{ + MaxFieldSectionSize: -1, + Datagram: pointer(false), + ExtendedConnect: pointer(false), + }, + expected: map[string]any{ + "frame_type": "settings", + "settings": []map[string]any{ + {"name": "settings_h3_datagram", "value": false}, + {"name": "settings_enable_connect_protocol", "value": false}, + }, + }, + }, + { + name: "unknowns", + // Only test a single unknown setting. + // Testing multiple unknown settings doesn't add a lot of value, + // and would require us to deal with non-deterministic map iteration order. + frame: SettingsFrame{ + MaxFieldSectionSize: -1, + Other: map[uint64]uint64{0xdead: 0xbeef}, + }, + expected: map[string]any{ + "frame_type": "settings", + "settings": []map[string]any{{ + "name": "unknown", + "name_bytes": float64(0xdead), + "value": float64(0xbeef), + }}, + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + check(t, tc.frame, tc.expected) + }) + } +} + +func TestPushPromiseFrame(t *testing.T) { + check(t, PushPromiseFrame{}, map[string]any{ + "frame_type": "push_promise", + }) +} + +func TestCancelPushFrame(t *testing.T) { + check(t, CancelPushFrame{}, map[string]any{ + "frame_type": "cancel_push", + }) +} + +func TestMaxPushIDFrame(t *testing.T) { + check(t, MaxPushIDFrame{}, map[string]any{ + "frame_type": "max_push_id", + }) +} + +func TestReservedFrame(t *testing.T) { + check(t, ReservedFrame{Type: 0x1f}, map[string]any{ + "frame_type": "reserved", + "frame_type_bytes": 0x1f, + }) +} + +func TestUnknownFrame(t *testing.T) { + check(t, UnknownFrame{Type: 0x2a}, map[string]any{ + "frame_type": "unknown", + "frame_type_bytes": 0x2a, + }) +} diff --git a/third_party/quic-go/http3/qlog/qlog_dir.go b/third_party/quic-go/http3/qlog/qlog_dir.go new file mode 100644 index 0000000..f74ec0a --- /dev/null +++ b/third_party/quic-go/http3/qlog/qlog_dir.go @@ -0,0 +1,15 @@ +package qlog + +import ( + "context" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +const EventSchema = "urn:ietf:params:qlog:events:http3-12" + +func DefaultConnectionTracer(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + return qlog.DefaultConnectionTracerWithSchemas(ctx, isClient, connID, []string{qlog.EventSchema, EventSchema}) +} diff --git a/third_party/quic-go/http3/qlog/qlog_dir_test.go b/third_party/quic-go/http3/qlog/qlog_dir_test.go new file mode 100644 index 0000000..fae09c7 --- /dev/null +++ b/third_party/quic-go/http3/qlog/qlog_dir_test.go @@ -0,0 +1,41 @@ +package qlog + +import ( + "context" + "os" + "path/filepath" + "testing" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlog" + "github.com/stretchr/testify/require" +) + +func TestQLOGDIRSet(t *testing.T) { + tmpDir := t.TempDir() + + connID := quic.ConnectionIDFromBytes([]byte{1, 2, 3, 4}) + qlogDir := filepath.Join(tmpDir, "qlogs") + t.Setenv("QLOGDIR", qlogDir) + + tracer := DefaultConnectionTracer(context.Background(), true, connID) + require.NotNil(t, tracer) + + // adding and closing a producer makes the tracer close the file + recorder := tracer.AddProducer() + recorder.Close() + + _, err := os.Stat(qlogDir) + qlogDirCreated := !os.IsNotExist(err) + require.True(t, qlogDirCreated) + + entries, err := os.ReadDir(qlogDir) + require.NoError(t, err) + require.Len(t, entries, 1) + + data, err := os.ReadFile(filepath.Join(qlogDir, entries[0].Name())) + require.NoError(t, err) + + require.Contains(t, string(data), EventSchema) + require.Contains(t, string(data), qlog.EventSchema) +} diff --git a/third_party/quic-go/http3/request_writer.go b/third_party/quic-go/http3/request_writer.go new file mode 100644 index 0000000..acd4e65 --- /dev/null +++ b/third_party/quic-go/http3/request_writer.go @@ -0,0 +1,324 @@ +package http3 + +import ( + "bytes" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptrace" + "strconv" + "strings" + "sync" + + "golang.org/x/net/http/httpguts" + "golang.org/x/net/http2/hpack" + "golang.org/x/net/idna" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/quic-go/qpack" +) + +const bodyCopyBufferSize = 8 * 1024 + +type requestWriter struct { + mutex sync.Mutex + encoder *qpack.Encoder + headerBuf *bytes.Buffer +} + +func newRequestWriter() *requestWriter { + headerBuf := &bytes.Buffer{} + encoder := qpack.NewEncoder(headerBuf) + return &requestWriter{ + encoder: encoder, + headerBuf: headerBuf, + } +} + +func (w *requestWriter) WriteRequestHeader(wr io.Writer, req *http.Request, gzip bool, streamID quic.StreamID, qlogger qlogwriter.Recorder) error { + buf := &bytes.Buffer{} + if err := w.writeHeaders(buf, req, gzip, streamID, qlogger); err != nil { + return err + } + if _, err := wr.Write(buf.Bytes()); err != nil { + return err + } + trace := httptrace.ContextClientTrace(req.Context()) + traceWroteHeaders(trace) + return nil +} + +func (w *requestWriter) writeHeaders(wr io.Writer, req *http.Request, gzip bool, streamID quic.StreamID, qlogger qlogwriter.Recorder) error { + w.mutex.Lock() + defer w.mutex.Unlock() + defer w.encoder.Close() + defer w.headerBuf.Reset() + + var trailers string + if len(req.Trailer) > 0 { + keys := make([]string, 0, len(req.Trailer)) + for k := range req.Trailer { + if httpguts.ValidTrailerHeader(k) { + keys = append(keys, k) + } + } + trailers = strings.Join(keys, ", ") + } + + headerFields, err := w.encodeHeaders(req, gzip, trailers, actualContentLength(req), qlogger != nil) + if err != nil { + return err + } + + b := make([]byte, 0, 128) + b = (&headersFrame{Length: uint64(w.headerBuf.Len())}).Append(b) + if qlogger != nil { + qlogCreatedHeadersFrame(qlogger, streamID, len(b)+w.headerBuf.Len(), w.headerBuf.Len(), headerFields) + } + if _, err := wr.Write(b); err != nil { + return err + } + _, err = wr.Write(w.headerBuf.Bytes()) + return err +} + +func isExtendedConnectRequest(req *http.Request) bool { + return req.Method == http.MethodConnect && req.Proto != "" && req.Proto != "HTTP/1.1" +} + +// copied from net/transport.go +// Modified to support Extended CONNECT: +// Contrary to what the godoc for the http.Request says, +// we do respect the Proto field if the method is CONNECT. +// +// The returned header fields are only set if doQlog is true. +func (w *requestWriter) encodeHeaders(req *http.Request, addGzipHeader bool, trailers string, contentLength int64, doQlog bool) ([]qlog.HeaderField, error) { + host := req.Host + if host == "" { + host = req.URL.Host + } + host, err := httpguts.PunycodeHostPort(host) + if err != nil { + return nil, err + } + if !httpguts.ValidHostHeader(host) { + return nil, errors.New("http3: invalid Host header") + } + + // http.NewRequest sets this field to HTTP/1.1 + isExtendedConnect := isExtendedConnectRequest(req) + if isExtendedConnect && !validExtendedConnectProtocol(req.Proto) { + return nil, fmt.Errorf("invalid request :protocol %q", req.Proto) + } + + var path string + if req.Method != http.MethodConnect || isExtendedConnect { + path = req.URL.RequestURI() + if !validPseudoPath(path) { + orig := path + path = strings.TrimPrefix(path, req.URL.Scheme+"://"+host) + if !validPseudoPath(path) { + if req.URL.Opaque != "" { + return nil, fmt.Errorf("invalid request :path %q from URL.Opaque = %q", orig, req.URL.Opaque) + } else { + return nil, fmt.Errorf("invalid request :path %q", orig) + } + } + } + } + + // Check for any invalid headers and return an error before we + // potentially pollute our hpack state. (We want to be able to + // continue to reuse the hpack encoder for future requests) + for k, vv := range req.Header { + if !httpguts.ValidHeaderFieldName(k) { + return nil, fmt.Errorf("invalid HTTP header name %q", k) + } + for _, v := range vv { + if !httpguts.ValidHeaderFieldValue(v) { + return nil, fmt.Errorf("invalid HTTP header value for header %q", k) + } + } + } + + enumerateHeaders := func(f func(name, value string)) { + // 8.1.2.3 Request Pseudo-Header Fields + // The :path pseudo-header field includes the path and query parts of the + // target URI (the path-absolute production and optionally a '?' character + // followed by the query production (see Sections 3.3 and 3.4 of + // [RFC3986]). + f(":authority", host) + f(":method", req.Method) + if req.Method != http.MethodConnect || isExtendedConnect { + f(":path", path) + f(":scheme", req.URL.Scheme) + } + if isExtendedConnect { + f(":protocol", req.Proto) + } + if trailers != "" { + f("trailer", trailers) + } + + var didUA bool + for k, vv := range req.Header { + if strings.EqualFold(k, "host") || strings.EqualFold(k, "content-length") { + // Host is :authority, already sent. + // Content-Length is automatic, set below. + continue + } else if strings.EqualFold(k, "connection") || strings.EqualFold(k, "proxy-connection") || + strings.EqualFold(k, "transfer-encoding") || strings.EqualFold(k, "upgrade") || + strings.EqualFold(k, "keep-alive") { + // Per 8.1.2.2 Connection-Specific Header + // Fields, don't send connection-specific + // fields. We have already checked if any + // are error-worthy so just ignore the rest. + continue + } else if strings.EqualFold(k, "user-agent") { + // Match Go's http1 behavior: at most one + // User-Agent. If set to nil or empty string, + // then omit it. Otherwise if not mentioned, + // include the default (below). + didUA = true + if len(vv) < 1 { + continue + } + vv = vv[:1] + if vv[0] == "" { + continue + } + + } + + for _, v := range vv { + f(k, v) + } + } + if shouldSendReqContentLength(req.Method, contentLength) { + f("content-length", strconv.FormatInt(contentLength, 10)) + } + if addGzipHeader { + f("accept-encoding", "gzip") + } + if !didUA { + f("user-agent", defaultUserAgent) + } + } + + // Do a first pass over the headers counting bytes to ensure + // we don't exceed cc.peerMaxHeaderListSize. This is done as a + // separate pass before encoding the headers to prevent + // modifying the hpack state. + hlSize := uint64(0) + enumerateHeaders(func(name, value string) { + hf := hpack.HeaderField{Name: name, Value: value} + hlSize += uint64(hf.Size()) + }) + + // TODO: check maximum header list size + // if hlSize > cc.peerMaxHeaderListSize { + // return errRequestHeaderListSize + // } + + trace := httptrace.ContextClientTrace(req.Context()) + traceHeaders := traceHasWroteHeaderField(trace) + + // Header list size is ok. Write the headers. + var headerFields []qlog.HeaderField + if doQlog { + headerFields = make([]qlog.HeaderField, 0, len(req.Header)) + } + enumerateHeaders(func(name, value string) { + name = strings.ToLower(name) + w.encoder.WriteField(qpack.HeaderField{Name: name, Value: value}) + if traceHeaders { + traceWroteHeaderField(trace, name, value) + } + if doQlog { + headerFields = append(headerFields, qlog.HeaderField{Name: name, Value: value}) + } + }) + + return headerFields, nil +} + +// authorityAddr returns a given authority (a host/IP, or host:port / ip:port) +// and returns a host:port. The port 443 is added if needed. +func authorityAddr(authority string) (addr string) { + host, port, err := net.SplitHostPort(authority) + if err != nil { // authority didn't have a port + port = "443" + host = authority + } + if a, err := idna.ToASCII(host); err == nil { + host = a + } + // IPv6 address literal, without a port: + if strings.HasPrefix(host, "[") && strings.HasSuffix(host, "]") { + return host + ":" + port + } + return net.JoinHostPort(host, port) +} + +// validPseudoPath reports whether v is a valid :path pseudo-header +// value. It must be either: +// +// *) a non-empty string starting with '/' +// *) the string '*', for OPTIONS requests. +// +// For now this is only used a quick check for deciding when to clean +// up Opaque URLs before sending requests from the Transport. +// See golang.org/issue/16847 +// +// We used to enforce that the path also didn't start with "//", but +// Google's GFE accepts such paths and Chrome sends them, so ignore +// that part of the spec. See golang.org/issue/19103. +func validPseudoPath(v string) bool { + return (len(v) > 0 && v[0] == '/') || v == "*" +} + +// actualContentLength returns a sanitized version of +// req.ContentLength, where 0 actually means zero (not unknown) and -1 +// means unknown. +func actualContentLength(req *http.Request) int64 { + if req.Body == nil { + return 0 + } + if req.ContentLength != 0 { + return req.ContentLength + } + return -1 +} + +// shouldSendReqContentLength reports whether the http2.Transport should send +// a "content-length" request header. This logic is basically a copy of the net/http +// transferWriter.shouldSendContentLength. +// The contentLength is the corrected contentLength (so 0 means actually 0, not unknown). +// -1 means unknown. +func shouldSendReqContentLength(method string, contentLength int64) bool { + if contentLength > 0 { + return true + } + if contentLength < 0 { + return false + } + // For zero bodies, whether we send a content-length depends on the method. + // It also kinda doesn't matter for http2 either way, with END_STREAM. + switch method { + case "POST", "PUT", "PATCH": + return true + default: + return false + } +} + +// WriteRequestTrailer writes HTTP trailers to the stream. +// It should be called after the request body has been fully written. +func (w *requestWriter) WriteRequestTrailer(wr io.Writer, req *http.Request, streamID quic.StreamID, qlogger qlogwriter.Recorder) error { + _, err := writeTrailers(wr, req.Trailer, streamID, qlogger) + return err +} diff --git a/third_party/quic-go/http3/request_writer_test.go b/third_party/quic-go/http3/request_writer_test.go new file mode 100644 index 0000000..7c57aed --- /dev/null +++ b/third_party/quic-go/http3/request_writer_test.go @@ -0,0 +1,175 @@ +package http3 + +import ( + "bytes" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func decodeRequest(t *testing.T, str io.Reader, streamID quic.StreamID, eventRecorder *events.Recorder) map[string]string { + t.Helper() + + r := io.LimitedReader{R: str, N: 1000} + fp := frameParser{r: &r} + frame, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &headersFrame{}, frame) + headersFrame := frame.(*headersFrame) + data := make([]byte, headersFrame.Length) + _, err = io.ReadFull(&r, data) + require.NoError(t, err) + hfs := decodeQpackHeaderFields(t, data) + values := make(map[string]string) + for _, hf := range hfs { + values[hf.Name] = hf.Value + } + + headerFields := make([]qlog.HeaderField, len(hfs)) + for i, hf := range hfs { + headerFields[i] = qlog.HeaderField{Name: hf.Name, Value: hf.Value} + } + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameCreated{ + StreamID: streamID, + Raw: qlog.RawInfo{ + Length: int(1000 - r.N), + PayloadLength: int(headersFrame.Length), + }, + Frame: qlog.Frame{Frame: qlog.HeadersFrame{HeaderFields: headerFields}}, + }, + }, + eventRecorder.Events(qlog.FrameCreated{}), + ) + + return values +} + +func TestRequestWriterGetRequestGzip(t *testing.T) { + t.Run("gzip", func(t *testing.T) { + testRequestWriterGzip(t, true) + }) + t.Run("no gzip", func(t *testing.T) { + testRequestWriterGzip(t, false) + }) +} + +func testRequestWriterGzip(t *testing.T, gzip bool) { + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net/index.html?foo=bar", nil) + req.AddCookie(&http.Cookie{Name: "foo", Value: "bar"}) + req.AddCookie(&http.Cookie{Name: "baz", Value: "lorem ipsum"}) + + rw := newRequestWriter() + var eventRecorder events.Recorder + buf := &bytes.Buffer{} + require.NoError(t, rw.WriteRequestHeader(buf, req, gzip, 42, &eventRecorder)) + headerFields := decodeRequest(t, buf, 42, &eventRecorder) + require.Equal(t, "quic-go.net", headerFields[":authority"]) + require.Equal(t, http.MethodGet, headerFields[":method"]) + require.Equal(t, "/index.html?foo=bar", headerFields[":path"]) + require.Equal(t, "https", headerFields[":scheme"]) + require.Equal(t, `foo=bar; baz="lorem ipsum"`, headerFields["cookie"]) + switch gzip { + case true: + require.Equal(t, "gzip", headerFields["accept-encoding"]) + case false: + require.NotContains(t, headerFields, "accept-encoding") + } +} + +func TestRequestWriterInvalidHostHeader(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net/index.html?foo=bar", nil) + req.Host = "foo@bar" // @ is invalid + rw := newRequestWriter() + require.EqualError(t, + rw.WriteRequestHeader(&bytes.Buffer{}, req, false, 0, nil), + "http3: invalid Host header", + ) +} + +func TestRequestWriterInvalidHeaderValue(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net", nil) + req.Header.Set("Authorization", "Bearer secret\x00") + err := newRequestWriter().WriteRequestHeader(&bytes.Buffer{}, req, false, 0, nil) + require.EqualError(t, err, `invalid HTTP header value for header "Authorization"`) +} + +func TestRequestWriterConnect(t *testing.T) { + // httptest.NewRequest does not properly support the CONNECT method + req, err := http.NewRequest(http.MethodConnect, "https://quic-go.net/", nil) + require.NoError(t, err) + rw := newRequestWriter() + buf := &bytes.Buffer{} + var eventRecorder events.Recorder + require.NoError(t, rw.WriteRequestHeader(buf, req, false, 1337, &eventRecorder)) + headerFields := decodeRequest(t, buf, 1337, &eventRecorder) + require.Equal(t, http.MethodConnect, headerFields[":method"]) + require.Equal(t, "quic-go.net", headerFields[":authority"]) + require.NotContains(t, headerFields, ":path") + require.NotContains(t, headerFields, ":scheme") + require.NotContains(t, headerFields, ":protocol") +} + +func TestRequestWriterExtendedConnect(t *testing.T) { + // httptest.NewRequest does not properly support the CONNECT method + req, err := http.NewRequest(http.MethodConnect, "https://quic-go.net/", nil) + require.NoError(t, err) + req.Proto = "webtransport" + rw := newRequestWriter() + buf := &bytes.Buffer{} + var eventRecorder events.Recorder + require.NoError(t, rw.WriteRequestHeader(buf, req, false, 1234, &eventRecorder)) + headerFields := decodeRequest(t, buf, 1234, &eventRecorder) + require.Equal(t, "quic-go.net", headerFields[":authority"]) + require.Equal(t, http.MethodConnect, headerFields[":method"]) + require.Equal(t, "/", headerFields[":path"]) + require.Equal(t, "https", headerFields[":scheme"]) + require.Equal(t, "webtransport", headerFields[":protocol"]) +} + +func TestRequestWriterExtendedConnectInvalidProtocol(t *testing.T) { + // httptest.NewRequest does not properly support the CONNECT method + req, err := http.NewRequest(http.MethodConnect, "https://quic-go.net/", nil) + require.NoError(t, err) + req.Proto = "HTTP/3.0" + rw := newRequestWriter() + require.EqualError(t, + rw.WriteRequestHeader(&bytes.Buffer{}, req, false, 0, nil), + `invalid request :protocol "HTTP/3.0"`, + ) +} + +func TestRequestWriterTrailers(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "https://quic-go.net/upload", nil) + req.Trailer = http.Header{ + "Trailer1": []string{"foo"}, + "Trailer2": []string{"bar"}, + "Content-Length": []string{"42"}, // Content-Length is not a valid trailer + } + + rw := newRequestWriter() + buf := &bytes.Buffer{} + require.NoError(t, rw.WriteRequestHeader(buf, req, false, 42, nil)) + headers := decodeHeader(t, buf) + require.Len(t, headers["trailer"], 1) + require.Contains(t, headers["trailer"][0], "Trailer1") + require.Contains(t, headers["trailer"][0], "Trailer2") + require.NotContains(t, headers["trailer"][0], "Content-Length") + + require.NoError(t, rw.WriteRequestTrailer(buf, req, 42, nil)) + + trailers := decodeHeader(t, buf) + require.Equal(t, map[string][]string{ + "trailer1": {"foo"}, + "trailer2": {"bar"}, + }, trailers) +} diff --git a/third_party/quic-go/http3/response_writer.go b/third_party/quic-go/http3/response_writer.go new file mode 100644 index 0000000..ea27e61 --- /dev/null +++ b/third_party/quic-go/http3/response_writer.go @@ -0,0 +1,372 @@ +package http3 + +import ( + "bytes" + "fmt" + "log/slog" + "net/http" + "net/textproto" + "strconv" + "strings" + "time" + + "github.com/apernet/quic-go/http3/qlog" + "github.com/quic-go/qpack" + + "golang.org/x/net/http/httpguts" +) + +// The HTTPStreamer allows taking over a HTTP/3 stream. The interface is implemented by the http.ResponseWriter. +// When a stream is taken over, it's the caller's responsibility to close the stream. +type HTTPStreamer interface { + HTTPStream() *Stream +} + +const maxSmallResponseSize = 4096 + +type responseWriter struct { + str *Stream + + conn *rawConn + header http.Header + trailers map[string]struct{} + buf []byte + status int // status code passed to WriteHeader + + // for responses smaller than maxSmallResponseSize, we buffer calls to Write, + // and automatically add the Content-Length header + smallResponseBuf []byte + + contentLen int64 // if handler set valid Content-Length header + numWritten int64 // bytes written + headerComplete bool // set once WriteHeader is called with a status code >= 200 + headerWritten bool // set once the response header has been serialized to the stream + isHead bool + trailerWritten bool // set once the response trailers has been serialized to the stream + + hijacked bool // set on HTTPStream is called + + logger *slog.Logger +} + +var ( + _ http.ResponseWriter = &responseWriter{} + _ http.Flusher = &responseWriter{} + _ Settingser = &responseWriter{} + _ HTTPStreamer = &responseWriter{} + // make sure that we implement (some of the) methods used by the http.ResponseController + _ interface { + SetReadDeadline(time.Time) error + SetWriteDeadline(time.Time) error + Flush() + FlushError() error + } = &responseWriter{} +) + +func newResponseWriter(str *Stream, conn *rawConn, isHead bool, logger *slog.Logger) *responseWriter { + return &responseWriter{ + str: str, + conn: conn, + header: http.Header{}, + buf: make([]byte, frameHeaderLen), + isHead: isHead, + logger: logger, + } +} + +func (w *responseWriter) Header() http.Header { + return w.header +} + +func (w *responseWriter) WriteHeader(status int) { + if w.headerComplete { + return + } + + // http status must be 3 digits + if status < 100 || status > 999 { + panic(fmt.Sprintf("invalid WriteHeader code %v", status)) + } + w.status = status + + // immediately write 1xx headers + if status < 200 { + w.writeHeader(status) + return + } + + // We're done with headers once we write a status >= 200. + w.headerComplete = true + // Add Date header. + // This is what the standard library does. + // Can be disabled by setting the Date header to nil. + if _, ok := w.header["Date"]; !ok { + w.header.Set("Date", time.Now().UTC().Format(http.TimeFormat)) + } + // Content-Length checking + // use ParseUint instead of ParseInt, as negative values are invalid + if clen := w.header.Get("Content-Length"); clen != "" { + if cl, err := strconv.ParseUint(clen, 10, 63); err == nil { + w.contentLen = int64(cl) + } else { + // emit a warning for malformed Content-Length and remove it + logger := w.logger + if logger == nil { + logger = slog.Default() + } + logger.Error("Malformed Content-Length", "value", clen) + w.header.Del("Content-Length") + } + } +} + +func (w *responseWriter) sniffContentType(p []byte) { + // If no content type, apply sniffing algorithm to body. + // We can't use `w.header.Get` here since if the Content-Type was set to nil, we shouldn't do sniffing. + _, haveType := w.header["Content-Type"] + + // If the Content-Encoding was set and is non-blank, we shouldn't sniff the body. + hasCE := w.header.Get("Content-Encoding") != "" + if !hasCE && !haveType && len(p) > 0 { + w.header.Set("Content-Type", http.DetectContentType(p)) + } +} + +func (w *responseWriter) Write(p []byte) (int, error) { + bodyAllowed := bodyAllowedForStatus(w.status) + if !w.headerComplete { + w.sniffContentType(p) + w.WriteHeader(http.StatusOK) + bodyAllowed = true + } + if !bodyAllowed { + return 0, http.ErrBodyNotAllowed + } + + w.numWritten += int64(len(p)) + if w.contentLen != 0 && w.numWritten > w.contentLen { + return 0, http.ErrContentLength + } + + if w.isHead { + return len(p), nil + } + + if !w.headerWritten { + // Buffer small responses. + // This allows us to automatically set the Content-Length field. + if len(w.smallResponseBuf)+len(p) < maxSmallResponseSize { + w.smallResponseBuf = append(w.smallResponseBuf, p...) + return len(p), nil + } + } + return w.doWrite(p) +} + +func (w *responseWriter) doWrite(p []byte) (int, error) { + if !w.headerWritten { + w.sniffContentType(w.smallResponseBuf) + if err := w.writeHeader(w.status); err != nil { + return 0, maybeReplaceError(err) + } + w.headerWritten = true + } + + l := uint64(len(w.smallResponseBuf) + len(p)) + if l == 0 { + return 0, nil + } + df := &dataFrame{Length: l} + w.buf = w.buf[:0] + w.buf = df.Append(w.buf) + if w.str.qlogger != nil { + w.str.qlogger.RecordEvent(qlog.FrameCreated{ + StreamID: w.str.StreamID(), + Raw: qlog.RawInfo{Length: len(w.buf) + int(l), PayloadLength: int(l)}, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }) + } + if _, err := w.str.writeUnframed(w.buf); err != nil { + return 0, maybeReplaceError(err) + } + if len(w.smallResponseBuf) > 0 { + if _, err := w.str.writeUnframed(w.smallResponseBuf); err != nil { + return 0, maybeReplaceError(err) + } + w.smallResponseBuf = nil + } + var n int + if len(p) > 0 { + var err error + n, err = w.str.writeUnframed(p) + if err != nil { + return n, maybeReplaceError(err) + } + } + return n, nil +} + +func (w *responseWriter) writeHeader(status int) error { + var headerFields []qlog.HeaderField // only used for qlog + var headers bytes.Buffer + enc := qpack.NewEncoder(&headers) + if err := enc.WriteField(qpack.HeaderField{Name: ":status", Value: strconv.Itoa(status)}); err != nil { + return err + } + if w.str.qlogger != nil { + headerFields = append(headerFields, qlog.HeaderField{Name: ":status", Value: strconv.Itoa(status)}) + } + + // Handle trailer fields + if vals, ok := w.header["Trailer"]; ok { + for _, val := range vals { + for trailer := range strings.SplitSeq(val, ",") { + // We need to convert to the canonical header key value here because this will be called when using + // headers.Add or headers.Set. + trailer = textproto.CanonicalMIMEHeaderKey(strings.TrimSpace(trailer)) + w.declareTrailer(trailer) + } + } + } + + for k, v := range w.header { + if _, excluded := w.trailers[k]; excluded { + continue + } + // Ignore "Trailer:" prefixed headers + if strings.HasPrefix(k, http.TrailerPrefix) { + continue + } + for index := range v { + name := strings.ToLower(k) + value := v[index] + if err := enc.WriteField(qpack.HeaderField{Name: name, Value: value}); err != nil { + return err + } + if w.str.qlogger != nil { + headerFields = append(headerFields, qlog.HeaderField{Name: name, Value: value}) + } + } + } + + buf := make([]byte, 0, frameHeaderLen+headers.Len()) + buf = (&headersFrame{Length: uint64(headers.Len())}).Append(buf) + buf = append(buf, headers.Bytes()...) + + if w.str.qlogger != nil { + qlogCreatedHeadersFrame(w.str.qlogger, w.str.StreamID(), len(buf), headers.Len(), headerFields) + } + + _, err := w.str.writeUnframed(buf) + return err +} + +func (w *responseWriter) FlushError() error { + if !w.headerComplete { + w.WriteHeader(http.StatusOK) + } + _, err := w.doWrite(nil) + return err +} + +func (w *responseWriter) flushTrailers() { + if w.trailerWritten { + return + } + if err := w.writeTrailers(); err != nil { + if w.logger != nil { + w.logger.Debug("could not write trailers", "error", err) + } + } +} + +func (w *responseWriter) Flush() { + if err := w.FlushError(); err != nil { + if w.logger != nil { + w.logger.Debug("could not flush to stream", "error", err) + } + } +} + +// declareTrailer adds a trailer to the trailer list, while also validating that the trailer has a +// valid name. +func (w *responseWriter) declareTrailer(k string) { + if !httpguts.ValidTrailerHeader(k) { + // Forbidden by RFC 9110, section 6.5.1. + if w.logger != nil { + w.logger.Debug("ignoring invalid trailer", slog.String("header", k)) + } + return + } + if w.trailers == nil { + w.trailers = make(map[string]struct{}) + } + w.trailers[k] = struct{}{} +} + +// writeTrailers will write trailers to the stream if there are any. +func (w *responseWriter) writeTrailers() error { + // promote headers added via "Trailer:" convention as trailers, these can be added after + // streaming the status/headers have been written. + for k := range w.header { + if strings.HasPrefix(k, http.TrailerPrefix) { + w.declareTrailer(k) + } + } + + if len(w.trailers) == 0 { + return nil + } + + trailers := make(http.Header, len(w.trailers)) + for trailer := range w.trailers { + if vals, ok := w.header[trailer]; ok { + trailers[strings.TrimPrefix(trailer, http.TrailerPrefix)] = vals + } + } + + written, err := writeTrailers(w.str.datagramStream, trailers, w.str.StreamID(), w.str.qlogger) + if written { + w.trailerWritten = true + } + return err +} + +func (w *responseWriter) HTTPStream() *Stream { + w.hijacked = true + w.Flush() + return w.str +} + +func (w *responseWriter) wasStreamHijacked() bool { return w.hijacked } + +func (w *responseWriter) ReceivedSettings() <-chan struct{} { + return w.conn.ReceivedSettings() +} + +func (w *responseWriter) Settings() *Settings { + return w.conn.Settings() +} + +func (w *responseWriter) SetReadDeadline(deadline time.Time) error { + return w.str.SetReadDeadline(deadline) +} + +func (w *responseWriter) SetWriteDeadline(deadline time.Time) error { + return w.str.SetWriteDeadline(deadline) +} + +// copied from http2/http2.go +// bodyAllowedForStatus reports whether a given response status code +// permits a body. See RFC 2616, section 4.4. +func bodyAllowedForStatus(status int) bool { + switch { + case status >= 100 && status <= 199: + return false + case status == http.StatusNoContent: + return false + case status == http.StatusNotModified: + return false + } + return true +} diff --git a/third_party/quic-go/http3/response_writer_test.go b/third_party/quic-go/http3/response_writer_test.go new file mode 100644 index 0000000..6e5041b --- /dev/null +++ b/third_party/quic-go/http3/response_writer_test.go @@ -0,0 +1,256 @@ +package http3 + +import ( + "bytes" + "io" + "log/slog" + "net/http" + "testing" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type testResponseWriter struct { + *responseWriter + eventRecorder *events.Recorder + buf *bytes.Buffer +} + +func (rw *testResponseWriter) DecodeHeaders(t *testing.T, idx int) map[string][]string { + t.Helper() + + rw.Flush() + rw.flushTrailers() + startLen := rw.buf.Len() + frame, err := (&frameParser{r: rw.buf}).ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &headersFrame{}, frame) + payloadLen := frame.(*headersFrame).Length + data := make([]byte, payloadLen) + headerFrameLen := startLen - rw.buf.Len() + len(data) + _, err = io.ReadFull(rw.buf, data) + require.NoError(t, err) + hfs := decodeQpackHeaderFields(t, data) + + // check that the decoded header fields are properly logged + require.GreaterOrEqual(t, len(rw.eventRecorder.Events(qlog.FrameCreated{})), idx+1) + require.IsType(t, qlog.HeadersFrame{}, rw.eventRecorder.Events()[idx].(qlog.FrameCreated).Frame.Frame) + ev := rw.eventRecorder.Events()[idx].(qlog.FrameCreated) + assert.Equal(t, quic.StreamID(42), ev.StreamID) + assert.Equal(t, headerFrameLen, ev.Raw.Length, "raw.Length") + assert.Equal(t, int(payloadLen), ev.Raw.PayloadLength, "raw.PayloadLength") + + fields := make(map[string][]string) + for _, p := range hfs { + fields[p.Name] = append(fields[p.Name], p.Value) + require.Contains(t, + ev.Frame.Frame.(qlog.HeadersFrame).HeaderFields, + qlog.HeaderField{Name: p.Name, Value: p.Value}, + ) + } + + return fields +} + +func (rw *testResponseWriter) DecodeBody(t *testing.T) []byte { + t.Helper() + + frame, err := (&frameParser{r: rw.buf}).ParseNext(nil) + if err == io.EOF { + return nil + } + require.NoError(t, err) + require.IsType(t, &dataFrame{}, frame) + body := make([]byte, frame.(*dataFrame).Length) + _, err = io.ReadFull(rw.buf, body) + require.NoError(t, err) + return body +} + +func newTestResponseWriter(t *testing.T) *testResponseWriter { + var eventRecorder events.Recorder + buf := &bytes.Buffer{} + mockCtrl := gomock.NewController(t) + str := NewMockDatagramStream(mockCtrl) + str.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + str.EXPECT().Write(gomock.Any()).DoAndReturn(buf.Write).AnyTimes() + str.EXPECT().SetReadDeadline(gomock.Any()).Return(nil).AnyTimes() + str.EXPECT().SetWriteDeadline(gomock.Any()).Return(nil).AnyTimes() + rw := newResponseWriter( + newStream(str, nil, nil, func(io.Reader, *headersFrame) error { return nil }, &eventRecorder), + nil, + false, + slog.Default(), + ) + return &testResponseWriter{ + responseWriter: rw, + eventRecorder: &eventRecorder, + buf: buf, + } +} + +func TestResponseWriterInvalidStatus(t *testing.T) { + rw := newTestResponseWriter(t) + require.Panics(t, func() { rw.WriteHeader(99) }) + require.Panics(t, func() { rw.WriteHeader(1000) }) +} + +func TestResponseWriterHeader(t *testing.T) { + rw := newTestResponseWriter(t) + rw.Header().Add("Content-Length", "42") + rw.WriteHeader(http.StatusTeapot) // 418 + // repeated WriteHeader calls are ignored + rw.WriteHeader(http.StatusInternalServerError) + + // set cookies + http.SetCookie(rw, &http.Cookie{Name: "foo", Value: "bar"}) + http.SetCookie(rw, &http.Cookie{Name: "baz", Value: "lorem ipsum"}) + // write some data + rw.Write([]byte("foobar")) + + fields := rw.DecodeHeaders(t, 0) + require.Equal(t, []string{"418"}, fields[":status"]) + require.Equal(t, []string{"42"}, fields["content-length"]) + require.Equal(t, + []string{"foo=bar", `baz="lorem ipsum"`}, + fields["set-cookie"], + ) + require.Equal(t, []byte("foobar"), rw.DecodeBody(t)) +} + +func TestResponseWriterDataWithoutHeader(t *testing.T) { + rw := newTestResponseWriter(t) + rw.Write([]byte("foobar")) + + fields := rw.DecodeHeaders(t, 0) + require.Equal(t, []string{"200"}, fields[":status"]) + require.Equal(t, []byte("foobar"), rw.DecodeBody(t)) +} + +func TestResponseWriterDataStatusWithoutBody(t *testing.T) { + rw := newTestResponseWriter(t) + rw.WriteHeader(http.StatusNotModified) + n, err := rw.Write([]byte("foobar")) + require.Zero(t, n) + require.ErrorIs(t, err, http.ErrBodyNotAllowed) + + fields := rw.DecodeHeaders(t, 0) + require.Equal(t, []string{"304"}, fields[":status"]) + require.Empty(t, rw.DecodeBody(t)) +} + +func TestResponseWriterContentLength(t *testing.T) { + rw := newTestResponseWriter(t) + rw.Header().Set("Content-Length", "6") + n, err := rw.Write([]byte("foobar")) + require.Equal(t, 6, n) + require.NoError(t, err) + + n, err = rw.Write([]byte{0x42}) + require.Zero(t, n) + require.ErrorIs(t, err, http.ErrContentLength) + + fields := rw.DecodeHeaders(t, 0) + require.Equal(t, []string{"200"}, fields[":status"]) + require.Equal(t, []string{"6"}, fields["content-length"]) + require.Equal(t, []byte("foobar"), rw.DecodeBody(t)) +} + +func TestResponseWriterContentTypeSniffing(t *testing.T) { + t.Run("no content type", func(t *testing.T) { + testContentTypeSniffing(t, map[string]string{}, "text/html; charset=utf-8") + }) + + t.Run("explicit content type", func(t *testing.T) { + testContentTypeSniffing(t, map[string]string{"Content-Type": "text/plain"}, "text/plain") + }) + + t.Run("with content encoding", func(t *testing.T) { + testContentTypeSniffing(t, map[string]string{"Content-Encoding": "gzip"}, "") + }) +} + +func testContentTypeSniffing(t *testing.T, hdrs map[string]string, expectedContentType string) { + rw := newTestResponseWriter(t) + for k, v := range hdrs { + rw.Header().Set(k, v) + } + rw.Write([]byte("")) + + fields := rw.DecodeHeaders(t, 0) + require.Equal(t, []string{"200"}, fields[":status"]) + if expectedContentType == "" { + require.NotContains(t, fields, "content-type") + } else { + require.Equal(t, []string{expectedContentType}, fields["content-type"]) + } +} + +func TestResponseWriterEarlyHints(t *testing.T) { + rw := newTestResponseWriter(t) + rw.Header().Add("Link", "; rel=preload; as=style") + rw.Header().Add("Link", "; rel=preload; as=script") + rw.WriteHeader(http.StatusEarlyHints) // status 103 + + n, err := rw.Write([]byte("foobar")) + require.Equal(t, 6, n) + require.NoError(t, err) + + // Early Hints must have been received + fields := rw.DecodeHeaders(t, 0) + require.Equal(t, 2, len(fields)) + require.Equal(t, []string{"103"}, fields[":status"]) + require.Equal(t, + []string{"; rel=preload; as=style", "; rel=preload; as=script"}, + fields["link"], + ) + + // headers sent in the informational response must also be included in the final response + fields = rw.DecodeHeaders(t, 1) + require.Equal(t, 4, len(fields)) + require.Equal(t, []string{"200"}, fields[":status"]) + require.Contains(t, fields, "date") + require.Contains(t, fields, "content-type") + require.Equal(t, + []string{"; rel=preload; as=style", "; rel=preload; as=script"}, + fields["link"], + ) + + require.Equal(t, []byte("foobar"), rw.DecodeBody(t)) +} + +func TestResponseWriterTrailers(t *testing.T) { + rw := newTestResponseWriter(t) + + rw.Header().Add("Trailer", "key, Content-Length") // Content-Length is not a valid trailer + n, err := rw.Write([]byte("foobar")) + require.Equal(t, 6, n) + require.NoError(t, err) + + // writeTrailers needs to be called after writing the full body + headers := rw.DecodeHeaders(t, 0) + require.Equal(t, []string{"key, Content-Length"}, headers["trailer"]) + require.NotContains(t, headers, "foo") + require.Equal(t, []byte("foobar"), rw.DecodeBody(t)) + + // headers set after writing the body are trailers + rw.Header().Set("key", "value") // announced trailer + rw.Header().Set("foo", "bar") // this trailer was not announced, and will therefore be ignored + rw.Header().Set(http.TrailerPrefix+"lorem", "ipsum") // unannounced trailer with trailer prefix + rw.Header().Set("Content-Length", "999") // invalid trailer, will be ignored + require.NoError(t, rw.writeTrailers()) + + trailers := rw.DecodeHeaders(t, 2) + require.Equal(t, []string{"value"}, trailers["key"]) + require.Equal(t, []string{"ipsum"}, trailers["lorem"]) + // trailers without the trailer prefix that were not announced are ignored + require.NotContains(t, trailers, "foo") + // invalid trailers are ignored + require.NotContains(t, trailers, "content-length") +} diff --git a/third_party/quic-go/http3/server.go b/third_party/quic-go/http3/server.go new file mode 100644 index 0000000..1b731a4 --- /dev/null +++ b/third_party/quic-go/http3/server.go @@ -0,0 +1,782 @@ +package http3 + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "log/slog" + "net" + "net/http" + "slices" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" +) + +// NextProtoH3 is the ALPN protocol negotiated during the TLS handshake, for QUIC v1 and v2. +const NextProtoH3 = "h3" + +// StreamType is the stream type of a unidirectional stream. +type StreamType uint64 + +const ( + streamTypeControlStream = 0 + streamTypePushStream = 1 + streamTypeQPACKEncoderStream = 2 + streamTypeQPACKDecoderStream = 3 +) + +// A QUICListener listens for incoming QUIC connections. +type QUICListener interface { + Accept(context.Context) (*quic.Conn, error) + Addr() net.Addr + io.Closer +} + +var _ QUICListener = &quic.EarlyListener{} + +// ConfigureTLSConfig creates a new tls.Config which can be used +// to create a quic.Listener meant for serving HTTP/3. +func ConfigureTLSConfig(tlsConf *tls.Config) *tls.Config { + // Workaround for https://github.com/golang/go/issues/60506. + // This initializes the session tickets _before_ cloning the config. + _, _ = tlsConf.DecryptTicket(nil, tls.ConnectionState{}) + config := tlsConf.Clone() + config.NextProtos = []string{NextProtoH3} + if gfc := config.GetConfigForClient; gfc != nil { + config.GetConfigForClient = func(ch *tls.ClientHelloInfo) (*tls.Config, error) { + conf, err := gfc(ch) + if conf == nil || err != nil { + return conf, err + } + return ConfigureTLSConfig(conf), nil + } + } + return config +} + +// contextKey is a value for use with context.WithValue. It's used as +// a pointer so it fits in an interface{} without allocation. +type contextKey struct { + name string +} + +func (k *contextKey) String() string { return "quic-go/http3 context value " + k.name } + +// ServerContextKey is a context key. It can be used in HTTP +// handlers with Context.Value to access the server that +// started the handler. The associated value will be of +// type *http3.Server. +var ServerContextKey = &contextKey{"http3-server"} + +// RemoteAddrContextKey is a context key. It can be used in +// HTTP handlers with Context.Value to access the remote +// address of the connection. The associated value will be of +// type net.Addr. +// +// Use this value instead of [http.Request.RemoteAddr] if you +// require access to the remote address of the connection rather +// than its string representation. +var RemoteAddrContextKey = &contextKey{"remote-addr"} + +// listener contains info about specific listener added with addListener +type listener struct { + ln *QUICListener + port int // 0 means that no info about port is available + + // if this listener was constructed by the application, it won't be closed when the server is closed + createdLocally bool +} + +// Server is a HTTP/3 server. +type Server struct { + // Addr optionally specifies the UDP address for the server to listen on, + // in the form "host:port". + // + // When used by ListenAndServe and ListenAndServeTLS methods, if empty, + // ":https" (port 443) is used. See net.Dial for details of the address + // format. + // + // Otherwise, if Port is not set and underlying QUIC listeners do not + // have valid port numbers, the port part is used in Alt-Svc headers set + // with SetQUICHeaders. + Addr string + + // Port is used in Alt-Svc response headers set with SetQUICHeaders. If + // needed Port can be manually set when the Server is created. + // + // This is useful when a Layer 4 firewall is redirecting UDP traffic and + // clients must use a port different from the port the Server is + // listening on. + Port int + + // TLSConfig provides a TLS configuration for use by server. It must be + // set for ListenAndServe and Serve methods. + TLSConfig *tls.Config + + // QUICConfig provides the parameters for QUIC connection created with Serve. + // If nil, it uses reasonable default values. + // + // Configured versions are also used in Alt-Svc response header set with SetQUICHeaders. + QUICConfig *quic.Config + + // Handler is the HTTP request handler to use. If not set, defaults to + // http.NotFound. + Handler http.Handler + + // EnableDatagrams enables support for HTTP/3 datagrams (RFC 9297). + // If set to true, QUICConfig.EnableDatagrams will be set. + EnableDatagrams bool + + // MaxHeaderBytes controls the maximum number of bytes the server will + // read parsing the request HEADERS frame. It does not limit the size of + // the request body. If zero or negative, http.DefaultMaxHeaderBytes is + // used. + MaxHeaderBytes int + + // AdditionalSettings specifies additional HTTP/3 settings. + // It is invalid to specify any settings defined by RFC 9114 (HTTP/3) and RFC 9297 (HTTP Datagrams). + AdditionalSettings map[uint64]uint64 + + // StreamDispatcher, when set, is called after peeking the first HTTP/3 frame type + // on a bidirectional stream. If handled is true, the stream is handed over to the + // application and won't be processed by HTTP/3. + StreamDispatcher func(FrameType, *quic.Stream, error) (handled bool, err error) + + // StreamAdmission, when set, runs synchronously immediately after accepting + // a bidirectional stream and before starting a handler goroutine or reading + // the first frame byte. Returning false resets the stream without spawning a + // handler. The returned release function, when non-nil, is called exactly + // once after dispatch or HTTP request handling completes. + // + // This AutoCAR extension allows a caller to apply process-wide admission and + // a first-byte deadline before StreamDispatcher's frame-type peek. + StreamAdmission func(*quic.Stream) (release func(), admitted bool) + + // IdleTimeout specifies how long until idle clients connection should be + // closed. Idle refers only to the HTTP/3 layer, activity at the QUIC layer + // like PING frames are not considered. + // If zero or negative, there is no timeout. + IdleTimeout time.Duration + + // ConnContext optionally specifies a function that modifies the context used for a new connection c. + // The provided ctx has a ServerContextKey value. + ConnContext func(ctx context.Context, c *quic.Conn) context.Context + + Logger *slog.Logger + + mutex sync.RWMutex + listeners []listener + + closed bool + closeCtx context.Context // canceled when the server is closed + closeCancel context.CancelFunc // cancels the closeCtx + graceCtx context.Context // canceled when the server is closed or gracefully closed + graceCancel context.CancelFunc // cancels the graceCtx + connCount atomic.Int64 + connHandlingDone chan struct{} + + altSvcHeader string +} + +// ListenAndServe listens on the UDP address s.Addr and calls s.Handler to handle HTTP/3 requests on incoming connections. +// +// If s.Addr is blank, ":https" is used. +func (s *Server) ListenAndServe() error { + ln, err := s.setupListenerForConn(s.TLSConfig, nil) + if err != nil { + return err + } + defer s.removeListener(ln) + + return s.serveListener(*ln) +} + +// ListenAndServeTLS listens on the UDP address s.Addr and calls s.Handler to handle HTTP/3 requests on incoming connections. +// +// If s.Addr is blank, ":https" is used. +func (s *Server) ListenAndServeTLS(certFile, keyFile string) error { + var err error + certs := make([]tls.Certificate, 1) + certs[0], err = tls.LoadX509KeyPair(certFile, keyFile) + if err != nil { + return err + } + // We currently only use the cert-related stuff from tls.Config, + // so we don't need to make a full copy. + ln, err := s.setupListenerForConn(&tls.Config{Certificates: certs}, nil) + if err != nil { + return err + } + defer s.removeListener(ln) + + return s.serveListener(*ln) +} + +// Serve an existing UDP connection. +// It is possible to reuse the same connection for outgoing connections. +// Closing the server does not close the connection. +func (s *Server) Serve(conn net.PacketConn) error { + ln, err := s.setupListenerForConn(s.TLSConfig, conn) + if err != nil { + return err + } + defer s.removeListener(ln) + + return s.serveListener(*ln) +} + +// init initializes the contexts used for shutting down the server. +// It must be called with the mutex held. +func (s *Server) init() { + if s.closeCtx == nil { + s.closeCtx, s.closeCancel = context.WithCancel(context.Background()) + s.graceCtx, s.graceCancel = context.WithCancel(s.closeCtx) + } + s.connHandlingDone = make(chan struct{}, 1) +} + +func (s *Server) decreaseConnCount() { + if s.connCount.Add(-1) == 0 && s.graceCtx.Err() != nil { + close(s.connHandlingDone) + } +} + +// ServeQUICConn serves a single QUIC connection. +func (s *Server) ServeQUICConn(conn *quic.Conn) error { + s.mutex.Lock() + if s.closed { + s.mutex.Unlock() + return http.ErrServerClosed + } + + s.init() + s.mutex.Unlock() + + s.connCount.Add(1) + defer s.decreaseConnCount() + + return s.handleConn(conn) +} + +// ServeListener serves an existing QUIC listener. +// Make sure you use http3.ConfigureTLSConfig to configure a tls.Config +// and use it to construct a http3-friendly QUIC listener. +// Closing the server does not close the listener. It is the application's responsibility to close them. +// ServeListener always returns a non-nil error. After Shutdown or Close, the returned error is http.ErrServerClosed. +func (s *Server) ServeListener(ln QUICListener) error { + s.mutex.Lock() + if err := s.addListener(&ln, false); err != nil { + s.mutex.Unlock() + return err + } + s.mutex.Unlock() + defer s.removeListener(&ln) + + return s.serveListener(ln) +} + +func (s *Server) serveListener(ln QUICListener) error { + for { + conn, err := ln.Accept(s.graceCtx) + // server closed + if errors.Is(err, quic.ErrServerClosed) || s.graceCtx.Err() != nil { + return http.ErrServerClosed + } + if err != nil { + return err + } + s.connCount.Add(1) + go func() { + defer s.decreaseConnCount() + if err := s.handleConn(conn); err != nil { + if s.Logger != nil { + s.Logger.Debug("handling connection failed", "error", err) + } + } + }() + } +} + +var errServerWithoutTLSConfig = errors.New("use of http3.Server without TLSConfig") + +func (s *Server) setupListenerForConn(tlsConf *tls.Config, conn net.PacketConn) (*QUICListener, error) { + if tlsConf == nil { + return nil, errServerWithoutTLSConfig + } + + baseConf := ConfigureTLSConfig(tlsConf) + quicConf := s.QUICConfig + if quicConf == nil { + quicConf = &quic.Config{Allow0RTT: true} + } else { + quicConf = s.QUICConfig.Clone() + } + if s.EnableDatagrams { + quicConf.EnableDatagrams = true + } + + s.mutex.Lock() + defer s.mutex.Unlock() + closed := s.closed + if closed { + return nil, http.ErrServerClosed + } + + var ln QUICListener + var err error + if conn == nil { + addr := s.Addr + if addr == "" { + addr = ":https" + } + ln, err = quic.ListenAddrEarly(addr, baseConf, quicConf) + } else { + ln, err = quic.ListenEarly(conn, baseConf, quicConf) + } + if err != nil { + return nil, err + } + if err := s.addListener(&ln, true); err != nil { + return nil, err + } + return &ln, nil +} + +func extractPort(addr string) (int, error) { + _, portStr, err := net.SplitHostPort(addr) + if err != nil { + return 0, err + } + + portInt, err := net.LookupPort("tcp", portStr) + if err != nil { + return 0, err + } + return portInt, nil +} + +func (s *Server) generateAltSvcHeader() { + if len(s.listeners) == 0 { + // Don't announce any ports since no one is listening for connections + s.altSvcHeader = "" + return + } + + // This code assumes that we will use protocol.SupportedVersions if no quic.Config is passed. + + var altSvc []string + addPort := func(port int) { + altSvc = append(altSvc, fmt.Sprintf(`%s=":%d"; ma=2592000`, NextProtoH3, port)) + } + + if s.Port != 0 { + // if Port is specified, we must use it instead of the + // listener addresses since there's a reason it's specified. + addPort(s.Port) + } else { + // if we have some listeners assigned, try to find ports + // which we can announce, otherwise nothing should be announced + validPortsFound := false + for _, info := range s.listeners { + if info.port != 0 { + addPort(info.port) + validPortsFound = true + } + } + if !validPortsFound { + if port, err := extractPort(s.Addr); err == nil { + addPort(port) + } + } + } + + s.altSvcHeader = strings.Join(altSvc, ",") +} + +func (s *Server) addListener(l *QUICListener, createdLocally bool) error { + if s.closed { + return http.ErrServerClosed + } + s.init() + + laddr := (*l).Addr() + if port, err := extractPort(laddr.String()); err == nil { + s.listeners = append(s.listeners, listener{ln: l, port: port, createdLocally: createdLocally}) + } else { + logger := s.Logger + if logger == nil { + logger = slog.Default() + } + logger.Error("Unable to extract port from listener, will not be announced using SetQUICHeaders", "local addr", laddr, "error", err) + s.listeners = append(s.listeners, listener{ln: l, port: 0, createdLocally: createdLocally}) + } + s.generateAltSvcHeader() + return nil +} + +func (s *Server) removeListener(l *QUICListener) { + s.mutex.Lock() + defer s.mutex.Unlock() + + s.listeners = slices.DeleteFunc(s.listeners, func(info listener) bool { + return info.ln == l + }) + s.generateAltSvcHeader() +} + +func (s *Server) NewRawServerConn(conn *quic.Conn) (*RawServerConn, error) { + hconn, _, _, err := s.newRawServerConn(conn) + if err != nil { + return nil, err + } + return hconn, nil +} + +func (s *Server) newRawServerConn(conn *quic.Conn) (*RawServerConn, *quic.SendStream, qlogwriter.Recorder, error) { + var qlogger qlogwriter.Recorder + if qlogTrace := conn.QlogTrace(); qlogTrace != nil && qlogTrace.SupportsSchemas(qlog.EventSchema) { + qlogger = qlogTrace.AddProducer() + } + connCtx := conn.Context() + connCtx = context.WithValue(connCtx, ServerContextKey, s) + connCtx = context.WithValue(connCtx, http.LocalAddrContextKey, conn.LocalAddr()) + connCtx = context.WithValue(connCtx, RemoteAddrContextKey, conn.RemoteAddr()) + if s.ConnContext != nil { + connCtx = s.ConnContext(connCtx, conn) + if connCtx == nil { + panic("http3: ConnContext returned nil") + } + } + hconn := newRawServerConn( + conn, + s.EnableDatagrams, + s.IdleTimeout, + qlogger, + s.Logger, + connCtx, + s.Handler, + s.maxHeaderBytes(), + ) + + // open the control stream and send a SETTINGS frame, it's also used to send a GOAWAY frame later + // when the server is gracefully closed + ctrlStr, err := hconn.openControlStream(&settingsFrame{ + MaxFieldSectionSize: int64(s.maxHeaderBytes()), + Datagram: s.EnableDatagrams, + ExtendedConnect: true, + Other: s.AdditionalSettings, + }) + if err != nil { + return nil, nil, nil, fmt.Errorf("opening the control stream failed: %w", err) + } + return hconn, ctrlStr, qlogger, nil +} + +// handleConn handles the HTTP/3 exchange on a QUIC connection. +// It blocks until all HTTP handlers for all streams have returned. +func (s *Server) handleConn(conn *quic.Conn) error { + hconn, ctrlStr, qlogger, err := s.newRawServerConn(conn) + if err != nil { + return err + } + + var wg sync.WaitGroup + wg.Go(func() { + for { + str, err := conn.AcceptUniStream(context.Background()) + if err != nil { + return + } + go hconn.HandleUnidirectionalStream(str) + } + }) + + var nextStreamID quic.StreamID + var handleErr error + var inGracefulShutdown bool + // Process all requests immediately. + // It's the client's responsibility to decide which requests are eligible for 0-RTT. + ctx := s.graceCtx + for { + // The context used here is: + // * before graceful shutdown: s.graceCtx + // * after graceful shutdown: s.closeCtx + // This allows us to keep accepting (and resetting) streams after graceful shutdown has started. + str, err := conn.AcceptStream(ctx) + if err != nil { + // the underlying connection was closed (by either side) + if conn.Context().Err() != nil { + var appErr *quic.ApplicationError + if !errors.As(err, &appErr) || appErr.ErrorCode != quic.ApplicationErrorCode(ErrCodeNoError) { + handleErr = fmt.Errorf("accepting stream failed: %w", err) + } + break + } + // server (not gracefully) closed, close the connection immediately + if s.closeCtx.Err() != nil { + hconn.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "") + handleErr = http.ErrServerClosed + break + } + inGracefulShutdown = s.graceCtx.Err() != nil + if !inGracefulShutdown { + var appErr *quic.ApplicationError + if !errors.As(err, &appErr) || appErr.ErrorCode != quic.ApplicationErrorCode(ErrCodeNoError) { + handleErr = fmt.Errorf("accepting stream failed: %w", err) + } + break + } + + // gracefully closed, send GOAWAY frame and wait for requests to complete or grace period to end + // new requests will be rejected and shouldn't be sent + if qlogger != nil { + qlogger.RecordEvent(qlog.FrameCreated{ + StreamID: ctrlStr.StreamID(), + Frame: qlog.Frame{Frame: qlog.GoAwayFrame{StreamID: nextStreamID}}, + }) + } + // Send the GOAWAY frame in a separate Goroutine. + // Sending might block if the peer didn't grant enough flow control credit. + // Write is guaranteed to return once the connection is closed. + wg.Go(func() { + _, _ = ctrlStr.Write((&goAwayFrame{StreamID: nextStreamID}).Append(nil)) + }) + ctx = s.closeCtx + continue + } + if inGracefulShutdown { + str.CancelRead(quic.StreamErrorCode(ErrCodeRequestRejected)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestRejected)) + continue + } + + nextStreamID = str.StreamID() + 4 + release := func() {} + if s.StreamAdmission != nil { + var admitted bool + release, admitted = s.StreamAdmission(str) + if !admitted { + str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + continue + } + if release == nil { + release = func() {} + } + } + wg.Go(func() { + defer release() + // HandleRequestStream will return once the request has been handled, + // or the underlying connection is closed. + if s.StreamDispatcher == nil { + hconn.HandleRequestStream(str) + return + } + frameType, err := quicvarint.Peek(str) + handled, dispatchErr := s.StreamDispatcher(FrameType(frameType), str, err) + if dispatchErr != nil { + str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + return + } + if handled { + return + } + hconn.HandleRequestStream(str) + }) + } + wg.Wait() + return handleErr +} + +func (s *Server) maxHeaderBytes() int { + if s.MaxHeaderBytes <= 0 { + return http.DefaultMaxHeaderBytes + } + return s.MaxHeaderBytes +} + +// Close the server immediately, aborting requests and sending CONNECTION_CLOSE frames to connected clients. +// Close in combination with ListenAndServe() (instead of Serve()) may race if it is called before a UDP socket is established. +// It is the caller's responsibility to close any connection passed to ServeQUICConn. +func (s *Server) Close() error { + s.mutex.Lock() + defer s.mutex.Unlock() + + s.closed = true + // server is never used + if s.closeCtx == nil { + return nil + } + s.closeCancel() + + var err error + for _, l := range s.listeners { + if l.createdLocally { + if cerr := (*l.ln).Close(); cerr != nil && err == nil { + err = cerr + } + } + } + if s.connCount.Load() == 0 { + return err + } + // wait for all connections to be closed + <-s.connHandlingDone + return err +} + +// Shutdown gracefully shuts down the server without interrupting any active connections. +// The server sends a GOAWAY frame first, then or for all running requests to complete. +// Shutdown in combination with ListenAndServe may race if it is called before a UDP socket is established. +// It is recommended to use Serve instead. +func (s *Server) Shutdown(ctx context.Context) error { + s.mutex.Lock() + s.closed = true + // server was never used + if s.closeCtx == nil { + s.mutex.Unlock() + return nil + } + s.graceCancel() + + // close all listeners + var closeErrs []error + for _, l := range s.listeners { + if l.createdLocally { + if err := (*l.ln).Close(); err != nil { + closeErrs = append(closeErrs, err) + } + } + } + s.mutex.Unlock() + if len(closeErrs) > 0 { + return errors.Join(closeErrs...) + } + + if s.connCount.Load() == 0 { + return s.Close() + } + select { + case <-s.connHandlingDone: // all connections were closed + // When receiving a GOAWAY frame, HTTP/3 clients are expected to close the connection + // once all requests were successfully handled... + return s.Close() + case <-ctx.Done(): + // ... however, clients handling long-lived requests (and misbehaving clients), + // might not do so before the context is cancelled. + // In this case, we close the server, which closes all existing connections + // (expect those passed to ServeQUICConn). + _ = s.Close() + return ctx.Err() + } +} + +// ErrNoAltSvcPort is the error returned by SetQUICHeaders when no port was found +// for Alt-Svc to announce. This can happen if listening on a PacketConn without a port +// (UNIX socket, for example) and no port is specified in Server.Port or Server.Addr. +var ErrNoAltSvcPort = errors.New("no port can be announced, specify it explicitly using Server.Port or Server.Addr") + +// SetQUICHeaders can be used to set the proper headers that announce that this server supports HTTP/3. +// The values set by default advertise all the ports the server is listening on, but can be +// changed to a specific port by setting Server.Port before launching the server. +// If no listener's Addr().String() returns an address with a valid port, Server.Addr will be used +// to extract the port, if specified. +// For example, a server launched using ListenAndServe on an address with port 443 would set: +// +// Alt-Svc: h3=":443"; ma=2592000 +func (s *Server) SetQUICHeaders(hdr http.Header) error { + s.mutex.RLock() + defer s.mutex.RUnlock() + + if s.altSvcHeader == "" { + return ErrNoAltSvcPort + } + // use the map directly to avoid constant canonicalization since the key is already canonicalized + hdr["Alt-Svc"] = append(hdr["Alt-Svc"], s.altSvcHeader) + return nil +} + +// ListenAndServeQUIC listens on the UDP network address addr and calls the +// handler for HTTP/3 requests on incoming connections. http.DefaultServeMux is +// used when handler is nil. +func ListenAndServeQUIC(addr, certFile, keyFile string, handler http.Handler) error { + server := &Server{ + Addr: addr, + Handler: handler, + } + return server.ListenAndServeTLS(certFile, keyFile) +} + +// ListenAndServeTLS listens on the given network address for both TLS/TCP and QUIC +// connections in parallel. It returns if one of the two returns an error. +// http.DefaultServeMux is used when handler is nil. +// The correct Alt-Svc headers for QUIC are set. +func ListenAndServeTLS(addr, certFile, keyFile string, handler http.Handler) error { + // Load certs + var err error + certs := make([]tls.Certificate, 1) + certs[0], err = tls.LoadX509KeyPair(certFile, keyFile) + if err != nil { + return err + } + // We currently only use the cert-related stuff from tls.Config, + // so we don't need to make a full copy. + config := &tls.Config{ + Certificates: certs, + } + + if addr == "" { + addr = ":https" + } + + // Open the listeners + udpAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return err + } + udpConn, err := net.ListenUDP("udp", udpAddr) + if err != nil { + return err + } + defer udpConn.Close() + + if handler == nil { + handler = http.DefaultServeMux + } + // Start the servers + quicServer := &Server{ + TLSConfig: config, + Handler: handler, + } + + hErr := make(chan error, 1) + qErr := make(chan error, 1) + go func() { + hErr <- http.ListenAndServeTLS(addr, certFile, keyFile, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + quicServer.SetQUICHeaders(w.Header()) + handler.ServeHTTP(w, r) + })) + }() + go func() { + qErr <- quicServer.Serve(udpConn) + }() + + select { + case err := <-hErr: + quicServer.Close() + return err + case err := <-qErr: + // Cannot close the HTTP server or wait for requests to complete properly :/ + return err + } +} diff --git a/third_party/quic-go/http3/server_conn.go b/third_party/quic-go/http3/server_conn.go new file mode 100644 index 0000000..d31922b --- /dev/null +++ b/third_party/quic-go/http3/server_conn.go @@ -0,0 +1,259 @@ +package http3 + +import ( + "context" + "errors" + "io" + "log/slog" + "net/http" + "runtime" + "strconv" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlogwriter" + "github.com/quic-go/qpack" +) + +// RawServerConn is an HTTP/3 server connection. +// It can be used for advanced use cases where the application wants to manage the QUIC connection lifecycle. +type RawServerConn struct { + rawConn rawConn + + idleTimeout time.Duration + idleTimer *time.Timer + + serverContext context.Context + requestHandler http.Handler + maxHeaderBytes int + + decoder *qpack.Decoder + + qlogger qlogwriter.Recorder + logger *slog.Logger +} + +func newRawServerConn( + conn *quic.Conn, + enableDatagrams bool, + idleTimeout time.Duration, + qlogger qlogwriter.Recorder, + logger *slog.Logger, + serverContext context.Context, + requestHandler http.Handler, + maxHeaderBytes int, +) *RawServerConn { + c := &RawServerConn{ + idleTimeout: idleTimeout, + serverContext: serverContext, + requestHandler: requestHandler, + maxHeaderBytes: maxHeaderBytes, + decoder: qpack.NewDecoder(), + qlogger: qlogger, + logger: logger, + } + c.rawConn = *newRawConn(conn, enableDatagrams, c.onStreamsEmpty, nil, qlogger, logger) + if idleTimeout > 0 { + c.idleTimer = time.AfterFunc(idleTimeout, func() { + conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "idle timeout") + }) + } + return c +} + +func (c *RawServerConn) onStreamsEmpty() { + if c.idleTimeout > 0 { + c.idleTimer.Reset(c.idleTimeout) + } +} + +// CloseWithError closes the connection with the given error code and message. +func (c *RawServerConn) CloseWithError(code quic.ApplicationErrorCode, msg string) error { + if c.idleTimer != nil { + c.idleTimer.Stop() + } + return c.rawConn.CloseWithError(code, msg) +} + +// HandleRequestStream handles an HTTP/3 request on a bidirectional request stream. +// The stream can either be obtained by calling AcceptStream on the underlying QUIC connection, +// or (internally) by using the server's stream accept loop. +func (c *RawServerConn) HandleRequestStream(str *quic.Stream) { + hstr := c.rawConn.TrackStream(str) + c.handleRequestStream(hstr) +} + +func (c *RawServerConn) requestMaxHeaderBytes() int { + if c.maxHeaderBytes <= 0 { + return http.DefaultMaxHeaderBytes + } + return c.maxHeaderBytes +} + +func (c *RawServerConn) openControlStream(settings *settingsFrame) (*quic.SendStream, error) { + return c.rawConn.openControlStream(settings) +} + +func (c *RawServerConn) handleRequestStream(str *stateTrackingStream) { + if c.idleTimeout > 0 { + // This only applies if the stream is the first active stream, + // but it's ok to stop a stopped timer. + c.idleTimer.Stop() + } + + conn := &c.rawConn + qlogger := c.qlogger + decoder := c.decoder + connCtx := c.serverContext + maxHeaderBytes := c.requestMaxHeaderBytes() + + fp := &frameParser{closeConn: conn.CloseWithError, r: str, streamID: str.StreamID()} + frame, err := fp.ParseNext(qlogger) + if err != nil { + str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + return + } + hf, ok := frame.(*headersFrame) + if !ok { + conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "expected first frame to be a HEADERS frame") + return + } + if hf.Length > uint64(maxHeaderBytes) { + maybeQlogInvalidHeadersFrame(qlogger, str.StreamID(), hf.Length) + // stop the client from sending more data + str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + // send a 431 Response (Request Header Fields Too Large) + c.rejectWithHeaderFieldsTooLarge(str) + return + } + headerBlock := make([]byte, hf.Length) + if _, err := io.ReadFull(str, headerBlock); err != nil { + maybeQlogInvalidHeadersFrame(qlogger, str.StreamID(), hf.Length) + str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + return + } + decodeFn := decoder.Decode(headerBlock) + var hfs []qpack.HeaderField + if qlogger != nil { + hfs = make([]qpack.HeaderField, 0, 16) + } + req, err := requestFromHeaders(decodeFn, maxHeaderBytes, &hfs) + if qlogger != nil { + qlogParsedHeadersFrame(qlogger, str.StreamID(), hf, hfs) + } + if err != nil { + if errors.Is(err, errHeaderTooLarge) { + // stop the client from sending more data + str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad)) + // send a 431 Response (Request Header Fields Too Large) + c.rejectWithHeaderFieldsTooLarge(str) + return + } + + errCode := ErrCodeMessageError + var qpackErr *qpackError + if errors.As(err, &qpackErr) { + errCode = ErrCodeQPACKDecompressionFailed + } + str.CancelRead(quic.StreamErrorCode(errCode)) + str.CancelWrite(quic.StreamErrorCode(errCode)) + return + } + + connState := conn.ConnectionState().TLS + req.TLS = &connState + req.RemoteAddr = conn.RemoteAddr().String() + + // Check that the client doesn't send more data in DATA frames than indicated by the Content-Length header (if set). + // See section 4.1.2 of RFC 9114. + contentLength := int64(-1) + if _, ok := req.Header["Content-Length"]; ok && req.ContentLength >= 0 { + contentLength = req.ContentLength + } + hstr := newStream(str, conn, nil, func(r io.Reader, hf *headersFrame) error { + trailers, err := decodeTrailers(r, hf, maxHeaderBytes, decoder, qlogger, str.StreamID()) + if err != nil { + return err + } + req.Trailer = trailers + return nil + }, qlogger) + body := newRequestBody(hstr, contentLength, connCtx, conn.ReceivedSettings(), conn.Settings) + req.Body = body + + if c.logger != nil { + c.logger.Debug("handling request", "method", req.Method, "host", req.Host, "uri", req.RequestURI) + } + + ctx, cancel := context.WithCancel(connCtx) + req = req.WithContext(ctx) + context.AfterFunc(str.Context(), cancel) + + r := newResponseWriter(hstr, conn, req.Method == http.MethodHead, c.logger) + handler := c.requestHandler + if handler == nil { + handler = http.DefaultServeMux + } + + // It's the client's responsibility to decide which requests are eligible for 0-RTT. + var panicked bool + func() { + defer func() { + if p := recover(); p != nil { + panicked = true + if p == http.ErrAbortHandler { + return + } + // Copied from net/http/server.go + const size = 64 << 10 + buf := make([]byte, size) + buf = buf[:runtime.Stack(buf, false)] + logger := c.logger + if logger == nil { + logger = slog.Default() + } + logger.Error("http3: panic serving", "arg", p, "trace", string(buf)) + } + }() + handler.ServeHTTP(r, req) + }() + + if r.wasStreamHijacked() { + return + } + + // abort the stream when there is a panic + if panicked { + str.CancelRead(quic.StreamErrorCode(ErrCodeInternalError)) + str.CancelWrite(quic.StreamErrorCode(ErrCodeInternalError)) + return + } + + // response not written to the client yet, set Content-Length + if !r.headerWritten { + if _, haveCL := r.header["Content-Length"]; !haveCL { + r.header.Set("Content-Length", strconv.FormatInt(r.numWritten, 10)) + } + } + r.Flush() + r.flushTrailers() + + // If the EOF was read by the handler, CancelRead() is a no-op. + str.CancelRead(quic.StreamErrorCode(ErrCodeNoError)) + str.Close() +} + +func (c *RawServerConn) rejectWithHeaderFieldsTooLarge(str *stateTrackingStream) { + hstr := newStream(str, &c.rawConn, nil, nil, c.qlogger) + defer hstr.Close() + r := newResponseWriter(hstr, &c.rawConn, false, c.logger) + r.WriteHeader(http.StatusRequestHeaderFieldsTooLarge) + r.Flush() +} + +// HandleUnidirectionalStream handles an incoming unidirectional stream. +func (c *RawServerConn) HandleUnidirectionalStream(str *quic.ReceiveStream) { + c.rawConn.handleUnidirectionalStream(str, true) +} diff --git a/third_party/quic-go/http3/server_test.go b/third_party/quic-go/http3/server_test.go new file mode 100644 index 0000000..16aff36 --- /dev/null +++ b/third_party/quic-go/http3/server_test.go @@ -0,0 +1,885 @@ +package http3 + +import ( + "bytes" + "context" + "crypto/tls" + "fmt" + "io" + "log/slog" + "net" + "net/http" + "net/http/httptest" + "runtime" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/internal/testdata" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func TestConfigureTLSConfig(t *testing.T) { + t.Run("basic config", func(t *testing.T) { + conf := ConfigureTLSConfig(&tls.Config{}) + require.Equal(t, conf.NextProtos, []string{NextProtoH3}) + }) + + t.Run("ALPN set", func(t *testing.T) { + conf := ConfigureTLSConfig(&tls.Config{NextProtos: []string{"foo", "bar"}}) + require.Equal(t, []string{NextProtoH3}, conf.NextProtos) + }) + + // for configs that define GetConfigForClient, the ALPN is set to h3 + t.Run("GetConfigForClient", func(t *testing.T) { + staticConf := &tls.Config{NextProtos: []string{"foo", "bar"}} + conf := ConfigureTLSConfig(&tls.Config{ + GetConfigForClient: func(*tls.ClientHelloInfo) (*tls.Config, error) { + return staticConf, nil + }, + }) + innerConf, err := conf.GetConfigForClient(&tls.ClientHelloInfo{ServerName: "example.com"}) + require.NoError(t, err) + require.NotNil(t, innerConf) + require.Equal(t, []string{NextProtoH3}, innerConf.NextProtos) + // make sure the original config was not modified + require.Equal(t, []string{"foo", "bar"}, staticConf.NextProtos) + }) + + // GetConfigForClient might return a nil tls.Config + t.Run("GetConfigForClient returns nil", func(t *testing.T) { + conf := ConfigureTLSConfig(&tls.Config{ + GetConfigForClient: func(*tls.ClientHelloInfo) (*tls.Config, error) { + return nil, nil + }, + }) + innerConf, err := conf.GetConfigForClient(&tls.ClientHelloInfo{ServerName: "example.com"}) + require.NoError(t, err) + require.Nil(t, innerConf) + }) +} + +func TestServerSettings(t *testing.T) { + t.Run("enable datagrams", func(t *testing.T) { + testServerSettings(t, true, nil) + }) + t.Run("additional settings", func(t *testing.T) { + testServerSettings(t, false, map[uint64]uint64{13: 37}) + }) +} + +func testServerSettings(t *testing.T, enableDatagrams bool, other map[uint64]uint64) { + s := Server{ + EnableDatagrams: enableDatagrams, + AdditionalSettings: other, + } + s.init() + + testDone := make(chan struct{}) + defer close(testDone) + + clientConn, serverConn := newConnPair(t) + go s.handleConn(serverConn) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + settingsStr, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + + settingsStr.SetReadDeadline(time.Now().Add(time.Second)) + b := make([]byte, 1024) + n, err := settingsStr.Read(b) + require.NoError(t, err) + b = b[:n] + + typ, l, err := quicvarint.Parse(b) + require.NoError(t, err) + require.EqualValues(t, streamTypeControlStream, typ) + fp := (&frameParser{r: bytes.NewReader(b[l:])}) + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f) + settingsFrame := f.(*settingsFrame) + // Extended CONNECT is always supported + require.True(t, settingsFrame.ExtendedConnect) + require.Equal(t, settingsFrame.Datagram, enableDatagrams) + require.Equal(t, settingsFrame.Other, other) +} + +func TestServerRequestHandling(t *testing.T) { + t.Run("200 with an empty handler", func(t *testing.T) { + var eventRecorder events.Recorder + hfs, body := testServerRequestHandling(t, + func(w http.ResponseWriter, r *http.Request) {}, + httptest.NewRequest(http.MethodGet, "https://www.example.com", nil), + &eventRecorder, + ) + require.Equal(t, hfs[":status"], []string{"200"}) + require.Empty(t, body) + + require.Len(t, eventRecorder.Events(qlog.FrameParsed{}), 1) + require.IsType(t, qlog.HeadersFrame{}, eventRecorder.Events(qlog.FrameParsed{})[0].(qlog.FrameParsed).Frame.Frame) + fp := eventRecorder.Events(qlog.FrameParsed{})[0].(qlog.FrameParsed) + require.Equal(t, quic.StreamID(0), fp.StreamID) + require.NotZero(t, fp.Raw.PayloadLength) + require.Contains(t, fp.Frame.Frame.(qlog.HeadersFrame).HeaderFields, qlog.HeaderField{Name: ":method", Value: "GET"}) + require.Contains(t, fp.Frame.Frame.(qlog.HeadersFrame).HeaderFields, qlog.HeaderField{Name: ":authority", Value: "www.example.com"}) + + events := filterQlogEventsForFrame(eventRecorder.Events(qlog.FrameCreated{}), qlog.HeadersFrame{}) + require.Len(t, events, 1) + fc := events[0].(qlog.FrameCreated) + require.Equal(t, quic.StreamID(0), fp.StreamID) + require.NotZero(t, fc.Raw.PayloadLength) + require.Contains(t, fc.Frame.Frame.(qlog.HeadersFrame).HeaderFields, qlog.HeaderField{Name: ":status", Value: "200"}) + }) + + t.Run("content-length", func(t *testing.T) { + hfs, body := testServerRequestHandling(t, + func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusTeapot) + w.Write([]byte("foobar")) + }, + httptest.NewRequest(http.MethodGet, "https://www.example.com", nil), + nil, + ) + require.Equal(t, hfs[":status"], []string{"418"}) + require.Equal(t, hfs["content-length"], []string{"6"}) + require.Equal(t, body, []byte("foobar")) + }) + + t.Run("no content-length when flushed", func(t *testing.T) { + hfs, body := testServerRequestHandling(t, + func(w http.ResponseWriter, r *http.Request) { + w.Write([]byte("foo")) + w.(http.Flusher).Flush() + w.Write([]byte("bar")) + }, + httptest.NewRequest(http.MethodGet, "https://www.example.com", nil), + nil, + ) + require.Equal(t, hfs[":status"], []string{"200"}) + require.NotContains(t, hfs, "content-length") + require.Equal(t, body, []byte("foobar")) + }) + + t.Run("HEAD request", func(t *testing.T) { + hfs, body := testServerRequestHandling(t, + func(w http.ResponseWriter, r *http.Request) { + w.Write([]byte("foobar")) + }, + httptest.NewRequest(http.MethodHead, "https://www.example.com", nil), + nil, + ) + require.Equal(t, hfs[":status"], []string{"200"}) + require.Empty(t, body) + }) + + t.Run("POST request", func(t *testing.T) { + hfs, body := testServerRequestHandling(t, + func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusTeapot) + data, _ := io.ReadAll(r.Body) + w.Write(data) + }, + httptest.NewRequest(http.MethodPost, "https://www.example.com", bytes.NewBuffer([]byte("foobar"))), + nil, + ) + require.Equal(t, hfs[":status"], []string{"418"}) + require.Equal(t, []byte("foobar"), body) + }) +} + +func testServerRequestHandling(t *testing.T, + handler http.HandlerFunc, + req *http.Request, + rec qlogwriter.Recorder, +) (responseHeaders map[string][]string, body []byte) { + clientConn, serverConn := newConnPair(t, withServerRecorder(rec)) + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = str.Write(encodeRequest(t, req)) + require.NoError(t, err) + require.NoError(t, str.Close()) + + s := &Server{Handler: handler} + go s.ServeQUICConn(serverConn) + + hfs := decodeHeader(t, str) + fp := frameParser{r: str} + var content []byte + for { + frame, err := fp.ParseNext(nil) + if err == io.EOF { + break + } + require.NoError(t, err) + require.IsType(t, &dataFrame{}, frame) + b := make([]byte, frame.(*dataFrame).Length) + _, err = io.ReadFull(str, b) + require.NoError(t, err) + content = append(content, b...) + } + return hfs, content +} + +func TestServerFirstFrameNotHeaders(t *testing.T) { + clientConn, serverConn := newConnPair(t) + str, err := clientConn.OpenStream() + require.NoError(t, err) + + var buf bytes.Buffer + buf.Write((&dataFrame{Length: 6}).Append(nil)) + buf.Write([]byte("foobar")) + _, err = str.Write(buf.Bytes()) + require.NoError(t, err) + require.NoError(t, str.Close()) + + s := &Server{} + go s.ServeQUICConn(serverConn) + + select { + case <-clientConn.Context().Done(): + err := context.Cause(clientConn.Context()) + var appErr *quic.ApplicationError + require.ErrorAs(t, err, &appErr) + require.Equal(t, quic.ApplicationErrorCode(ErrCodeFrameUnexpected), appErr.ErrorCode) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestStreamAdmissionRunsBeforeFramePeekAndReleases(t *testing.T) { + clientConn, serverConn := newConnPair(t) + admitted := make(chan struct{}) + released := make(chan struct{}) + dispatched := make(chan error, 1) + server := &Server{ + StreamAdmission: func(stream *quic.Stream) (func(), bool) { + close(admitted) + require.NoError(t, stream.SetReadDeadline(time.Now().Add(20*time.Millisecond))) + return func() { close(released) }, true + }, + StreamDispatcher: func(_ FrameType, _ *quic.Stream, err error) (bool, error) { + dispatched <- err + return true, nil + }, + } + go server.ServeQUICConn(serverConn) + + stream, err := clientConn.OpenStream() + require.NoError(t, err) + defer stream.Close() + // A QUIC stream isn't visible to the peer until it carries data. Send only + // the first byte of a two-byte varint so AcceptStream returns while the + // frame-type peek itself remains incomplete. + _, err = stream.Write([]byte{0x40}) + require.NoError(t, err) + select { + case <-admitted: + // Reaching admission while the frame varint is incomplete proves it + // precedes quicvarint.Peek. + case <-time.After(time.Second): + t.Fatal("stream was not admitted before the frame peek") + } + select { + case err := <-dispatched: + require.Error(t, err) + case <-time.After(time.Second): + t.Fatal("pre-peek read deadline did not interrupt dispatch") + } + select { + case <-released: + case <-time.After(time.Second): + t.Fatal("admission release callback was not called") + } +} + +func TestStreamAdmissionHoldsSlotUntilDispatcherReturns(t *testing.T) { + clientConn, serverConn := newConnPair(t) + slot := make(chan struct{}, 1) + dispatchStarted := make(chan struct{}) + unblockDispatch := make(chan struct{}) + rejected := make(chan struct{}) + released := make(chan struct{}) + server := &Server{ + StreamAdmission: func(_ *quic.Stream) (func(), bool) { + select { + case slot <- struct{}{}: + return func() { + <-slot + close(released) + }, true + default: + close(rejected) + return nil, false + } + }, + StreamDispatcher: func(_ FrameType, _ *quic.Stream, _ error) (bool, error) { + close(dispatchStarted) + <-unblockDispatch + return true, nil + }, + } + go server.ServeQUICConn(serverConn) + + first, err := clientConn.OpenStream() + require.NoError(t, err) + defer first.Close() + _, err = first.Write(quicvarint.Append(nil, 0x21)) + require.NoError(t, err) + select { + case <-dispatchStarted: + case <-time.After(time.Second): + t.Fatal("first stream was not dispatched") + } + + second, err := clientConn.OpenStream() + require.NoError(t, err) + defer second.Close() + _, _ = second.Write(quicvarint.Append(nil, 0x21)) + select { + case <-rejected: + case <-time.After(time.Second): + t.Fatal("second stream bypassed the occupied admission slot") + } + select { + case <-released: + t.Fatal("slot was released while the dispatcher was still running") + default: + } + + close(unblockDispatch) + select { + case <-released: + case <-time.After(time.Second): + t.Fatal("slot was not released after the dispatcher returned") + } +} + +func TestServerHandlerBodyNotRead(t *testing.T) { + t.Run("GET request with a body", func(t *testing.T) { + testServerHandlerBodyNotRead(t, + httptest.NewRequest(http.MethodGet, "https://www.example.com", bytes.NewBuffer([]byte("foobar"))), + func(w http.ResponseWriter, r *http.Request) {}, + ) + }) + + t.Run("POST body not read", func(t *testing.T) { + testServerHandlerBodyNotRead(t, + httptest.NewRequest(http.MethodPost, "https://www.example.com", bytes.NewBuffer([]byte("foobar"))), + func(w http.ResponseWriter, r *http.Request) {}, + ) + }) + + t.Run("POST request, with a replaced body", func(t *testing.T) { + testServerHandlerBodyNotRead(t, + httptest.NewRequest(http.MethodPost, "https://www.example.com", bytes.NewBuffer([]byte("foobar"))), + func(w http.ResponseWriter, r *http.Request) { + r.Body = struct { + io.Reader + io.Closer + }{} + }, + ) + }) +} + +func testServerHandlerBodyNotRead(t *testing.T, req *http.Request, handler http.HandlerFunc) { + clientConn, serverConn := newConnPair(t) + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = str.Write(encodeRequest(t, req)) + require.NoError(t, err) + + done := make(chan struct{}) + s := &Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer close(done) + handler(w, r) + }), + } + + go s.ServeQUICConn(serverConn) + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestServerStreamResetByClient(t *testing.T) { + clientConn, serverConn := newConnPair(t) + str, err := clientConn.OpenStream() + require.NoError(t, err) + str.CancelWrite(1337) + + var called bool + s := &Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true + }), + } + + go s.ServeQUICConn(serverConn) + + expectStreamReadReset(t, str, quic.StreamErrorCode(ErrCodeRequestIncomplete)) + require.False(t, called) +} + +func TestServerPanickingHandler(t *testing.T) { + t.Run("panicking handler", func(t *testing.T) { + logOutput := testServerPanickingHandler(t, func(w http.ResponseWriter, r *http.Request) { + panic("foobar") + }) + require.Contains(t, logOutput, "http3: panic serving") + require.Contains(t, logOutput, "foobar") + }) + + t.Run("http.ErrAbortHandler", func(t *testing.T) { + logOutput := testServerPanickingHandler(t, func(w http.ResponseWriter, r *http.Request) { + panic(http.ErrAbortHandler) + }) + require.NotContains(t, logOutput, "http3: panic serving") + require.NotContains(t, logOutput, "http.ErrAbortHandler") + }) +} + +func testServerPanickingHandler(t *testing.T, handler http.HandlerFunc) (logOutput string) { + clientConn, serverConn := newConnPair(t) + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = str.Write(encodeRequest(t, httptest.NewRequest(http.MethodHead, "https://www.example.com", nil))) + require.NoError(t, err) + require.NoError(t, str.Close()) + + var logBuf bytes.Buffer + s := &Server{ + Handler: handler, + Logger: slog.New(slog.NewTextHandler(&logBuf, nil)), + } + + go s.ServeQUICConn(serverConn) + + expectStreamReadReset(t, str, quic.StreamErrorCode(ErrCodeInternalError)) + s.Close() + + return logBuf.String() +} + +func TestServerRequestHeaderTooLarge(t *testing.T) { + t.Run("default value", func(t *testing.T) { + var eventRecorder events.Recorder + // use 2*DefaultMaxHeaderBytes here. qpack will compress the request, + // but the request will still end up larger than DefaultMaxHeaderBytes. + url := bytes.Repeat([]byte{'a'}, http.DefaultMaxHeaderBytes*2) + testServerRequestHeaderTooLarge(t, + httptest.NewRequest(http.MethodGet, "https://"+string(url), nil), + 0, + &eventRecorder, + ) + events := eventRecorder.Events(qlog.FrameParsed{}) + require.Len(t, events, 1) + require.Equal(t, qlog.HeadersFrame{}, events[0].(qlog.FrameParsed).Frame.Frame) + // The request is QPACK-compressed, so it will be smaller than 2*http.DefaultMaxHeaderBytes + require.Greater(t, events[0].(qlog.FrameParsed).Raw.PayloadLength, http.DefaultMaxHeaderBytes) + require.Less(t, events[0].(qlog.FrameParsed).Raw.PayloadLength, http.DefaultMaxHeaderBytes*2) + }) + + t.Run("custom value", func(t *testing.T) { + var eventRecorder events.Recorder + testServerRequestHeaderTooLarge(t, + httptest.NewRequest(http.MethodGet, "https://www.example.com", nil), + 20, + &eventRecorder, + ) + events := eventRecorder.Events(qlog.FrameParsed{}) + require.Len(t, events, 1) + require.Equal(t, qlog.HeadersFrame{}, events[0].(qlog.FrameParsed).Frame.Frame) + require.Greater(t, events[0].(qlog.FrameParsed).Raw.PayloadLength, 20) + require.Less(t, events[0].(qlog.FrameParsed).Raw.PayloadLength, 40) + }) +} + +func testServerRequestHeaderTooLarge(t *testing.T, req *http.Request, maxHeaderBytes int, rec qlogwriter.Recorder) { + var called bool + s := &Server{ + MaxHeaderBytes: maxHeaderBytes, + Handler: http.HandlerFunc(func(http.ResponseWriter, *http.Request) { called = true }), + } + s.init() + + clientConn, serverConn := newConnPair(t, withServerRecorder(rec)) + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = str.Write(encodeRequest(t, req)) + require.NoError(t, err) + + go s.ServeQUICConn(serverConn) + + hfs := decodeHeader(t, str) + require.Equal(t, []string{"431"}, hfs[":status"]) + expectStreamWriteReset(t, str, quic.StreamErrorCode(ErrCodeExcessiveLoad)) + require.False(t, called) +} + +func TestServerRequestContext(t *testing.T) { + clientConn, serverConn := newConnPair(t) + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = str.Write(encodeRequest(t, httptest.NewRequest(http.MethodHead, "https://www.example.com", nil))) + require.NoError(t, err) + + ctxChan := make(chan context.Context, 1) + block := make(chan struct{}) + s := &Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ctxChan <- r.Context() + <-block + }), + } + + go s.ServeQUICConn(serverConn) + + var requestContext context.Context + select { + case requestContext = <-ctxChan: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Equal(t, s, requestContext.Value(ServerContextKey)) + require.Equal(t, serverConn.LocalAddr(), requestContext.Value(http.LocalAddrContextKey)) + require.Equal(t, serverConn.RemoteAddr(), requestContext.Value(RemoteAddrContextKey)) + select { + case <-requestContext.Done(): + t.Fatal("request context was canceled") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + str.CancelRead(1337) + + select { + case <-requestContext.Done(): + case <-time.After(time.Second): + t.Fatal("timeout") + } + require.Equal(t, context.Canceled, requestContext.Err()) + close(block) +} + +func TestServerHTTPStreamHijacking(t *testing.T) { + clientConn, serverConn := newConnPair(t) + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = str.Write(encodeRequest(t, httptest.NewRequest(http.MethodHead, "https://www.example.com", nil))) + require.NoError(t, err) + require.NoError(t, str.Close()) + + s := &Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + str := w.(HTTPStreamer).HTTPStream() + str.Write([]byte("foobar")) + str.Close() + }), + } + go s.ServeQUICConn(serverConn) + + str.SetReadDeadline(time.Now().Add(time.Second)) + rsp, err := io.ReadAll(str) + require.NoError(t, err) + r := bytes.NewReader(rsp) + hfs := decodeHeader(t, r) + require.Equal(t, hfs[":status"], []string{"200"}) + fp := frameParser{r: r} + frame, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &dataFrame{}, frame) + dataFrame := frame.(*dataFrame) + require.Equal(t, uint64(6), dataFrame.Length) + data, err := io.ReadAll(r) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), data) +} + +func getAltSvc(s *Server) (string, bool) { + hdr := http.Header{} + s.SetQUICHeaders(hdr) + if altSvc, ok := hdr["Alt-Svc"]; ok { + return altSvc[0], true + } + return "", false +} + +func TestServerAltSvcFromListenersAndConns(t *testing.T) { + t.Run("default", func(t *testing.T) { + testServerAltSvcFromListenersAndConns(t, []quic.Version{}) + }) + t.Run("v1", func(t *testing.T) { + testServerAltSvcFromListenersAndConns(t, []quic.Version{quic.Version1}) + }) + t.Run("v1 and v2", func(t *testing.T) { + testServerAltSvcFromListenersAndConns(t, []quic.Version{quic.Version1, quic.Version2}) + }) +} + +func testServerAltSvcFromListenersAndConns(t *testing.T, versions []quic.Version) { + ln1, err := quic.ListenEarly(newUDPConnLocalhost(t), getTLSConfig(), nil) + require.NoError(t, err) + port1 := ln1.Addr().(*net.UDPAddr).Port + + s := &Server{ + Addr: ":1337", // will be ignored since we're using listeners + TLSConfig: getTLSConfig(), + QUICConfig: &quic.Config{Versions: versions}, + } + done1 := make(chan struct{}) + go func() { + defer close(done1) + s.ServeListener(ln1) + }() + time.Sleep(scaleDuration(10 * time.Millisecond)) + altSvc, ok := getAltSvc(s) + require.True(t, ok) + require.Equal(t, fmt.Sprintf(`h3=":%d"; ma=2592000`, port1), altSvc) + + udpConn := newUDPConnLocalhost(t) + port2 := udpConn.LocalAddr().(*net.UDPAddr).Port + done2 := make(chan struct{}) + go func() { + defer close(done2) + s.Serve(udpConn) + }() + time.Sleep(scaleDuration(10 * time.Millisecond)) + altSvc, ok = getAltSvc(s) + require.True(t, ok) + require.Equal(t, fmt.Sprintf(`h3=":%d"; ma=2592000,h3=":%d"; ma=2592000`, port1, port2), altSvc) + + // Close the first listener. + // This should remove the associated Alt-Svc entry. + require.NoError(t, ln1.Close()) + select { + case <-done1: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + altSvc, ok = getAltSvc(s) + require.True(t, ok) + require.Equal(t, fmt.Sprintf(`h3=":%d"; ma=2592000`, port2), altSvc) + + // Close the second listener. + // This should remove the Alt-Svc entry altogether. + require.NoError(t, udpConn.Close()) + select { + case <-done2: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + _, ok = getAltSvc(s) + require.False(t, ok) +} + +func TestServerAltSvcFromPort(t *testing.T) { + s := &Server{Port: 1337} + _, ok := getAltSvc(s) + require.False(t, ok) + + ln, err := quic.ListenEarly(newUDPConnLocalhost(t), getTLSConfig(), nil) + require.NoError(t, err) + done := make(chan struct{}) + go func() { + defer close(done) + s.ServeListener(ln) + }() + time.Sleep(scaleDuration(10 * time.Millisecond)) + + altSvc, ok := getAltSvc(s) + require.True(t, ok) + require.Equal(t, `h3=":1337"; ma=2592000`, altSvc) + + require.NoError(t, ln.Close()) + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + _, ok = getAltSvc(s) + require.False(t, ok) +} + +type unixSocketListener struct { + *quic.EarlyListener +} + +func (l *unixSocketListener) Addr() net.Addr { + return &net.UnixAddr{Net: "unix", Name: "/tmp/quic.sock"} +} + +func TestServerAltSvcFromUnixSocket(t *testing.T) { + t.Run("with Server.Addr not set", func(t *testing.T) { + _, ok := testServerAltSvcFromUnixSocket(t, "") + require.False(t, ok) + }) + + t.Run("with Server.Addr set", func(t *testing.T) { + altSvc, ok := testServerAltSvcFromUnixSocket(t, ":1337") + require.True(t, ok) + require.Equal(t, `h3=":1337"; ma=2592000`, altSvc) + }) +} + +func testServerAltSvcFromUnixSocket(t *testing.T, addr string) (altSvc string, ok bool) { + ln, err := quic.ListenEarly(newUDPConnLocalhost(t), testdata.GetTLSConfig(), nil) + require.NoError(t, err) + + var logBuf bytes.Buffer + s := &Server{ + Addr: addr, + Logger: slog.New(slog.NewTextHandler(&logBuf, nil)), + } + done := make(chan struct{}) + go func() { + defer close(done) + s.ServeListener(&unixSocketListener{EarlyListener: ln}) + }() + time.Sleep(scaleDuration(10 * time.Millisecond)) + + altSvc, ok = getAltSvc(s) + require.NoError(t, ln.Close()) + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + require.Contains(t, logBuf.String(), "Unable to extract port from listener, will not be announced using SetQUICHeaders") + return altSvc, ok +} + +func TestServerListenAndServeErrors(t *testing.T) { + require.EqualError(t, (&Server{}).ListenAndServe(), "use of http3.Server without TLSConfig") + s := &Server{ + Addr: ":123456", + TLSConfig: testdata.GetTLSConfig(), + } + require.ErrorContains(t, s.ListenAndServe(), "invalid port") +} + +func TestServerClosing(t *testing.T) { + s := &Server{TLSConfig: getTLSConfig()} + require.NoError(t, s.Close()) + require.NoError(t, s.Close()) // duplicate calls are ok + require.ErrorIs(t, s.ListenAndServe(), http.ErrServerClosed) + require.ErrorIs(t, s.ListenAndServeTLS(testdata.GetCertificatePaths()), http.ErrServerClosed) + require.ErrorIs(t, s.Serve(nil), http.ErrServerClosed) + require.ErrorIs(t, s.ServeListener(nil), http.ErrServerClosed) + require.ErrorIs(t, s.ServeQUICConn(nil), http.ErrServerClosed) +} + +func TestServerConcurrentServeAndClose(t *testing.T) { + addr, err := net.ResolveUDPAddr("udp", "localhost:0") + require.NoError(t, err) + c, err := net.ListenUDP("udp", addr) + require.NoError(t, err) + done := make(chan struct{}) + s := &Server{TLSConfig: testdata.GetTLSConfig()} + go func() { + defer close(done) + s.Serve(c) + }() + runtime.Gosched() + s.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestServerImmediateGracefulShutdown(t *testing.T) { + s := &Server{TLSConfig: testdata.GetTLSConfig()} + errChan := make(chan error, 1) + go func() { errChan <- s.Shutdown(context.Background()) }() + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestServerGracefulShutdown(t *testing.T) { + requestChan := make(chan struct{}, 1) + s := &Server{Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestChan <- struct{}{} + })} + + clientConn, serverConn := newConnPair(t) + go s.ServeQUICConn(serverConn) + + firstStream, err := clientConn.OpenStream() + require.NoError(t, err) + _, err = firstStream.Write(encodeRequest(t, httptest.NewRequest(http.MethodGet, "https://www.example.com", nil))) + require.NoError(t, err) + + select { + case <-requestChan: + case <-time.After(time.Second): + t.Fatal("timeout") + } + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + controlStr, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + typ, err := quicvarint.Read(quicvarint.NewReader(controlStr)) + require.NoError(t, err) + require.EqualValues(t, streamTypeControlStream, typ) + fp := &frameParser{r: controlStr} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + require.IsType(t, &settingsFrame{}, f) + + shutdownCtx, shutdownCancel := context.WithCancel(context.Background()) + errChan := make(chan error) + go func() { + errChan <- s.Shutdown(shutdownCtx) + }() + + f, err = fp.ParseNext(nil) + require.NoError(t, err) + require.Equal(t, &goAwayFrame{StreamID: 4}, f) + + select { + case <-errChan: + t.Fatal("didn't expect Shutdown to return") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + // all further streams are getting rejected + for range 3 { + str, err := clientConn.OpenStream() + require.NoError(t, err) + _, _ = str.Write(encodeRequest(t, httptest.NewRequest(http.MethodGet, "https://www.example.com", nil))) + expectStreamReadReset(t, str, quic.StreamErrorCode(ErrCodeRequestRejected)) + expectStreamWriteReset(t, str, quic.StreamErrorCode(ErrCodeRequestRejected)) + } + + // cancel the context passed to Shutdown + shutdownCancel() + + select { + case err := <-errChan: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} diff --git a/third_party/quic-go/http3/state_tracking_stream.go b/third_party/quic-go/http3/state_tracking_stream.go new file mode 100644 index 0000000..c7cf302 --- /dev/null +++ b/third_party/quic-go/http3/state_tracking_stream.go @@ -0,0 +1,173 @@ +package http3 + +import ( + "context" + "errors" + "os" + "sync" + + "github.com/apernet/quic-go" +) + +const streamDatagramQueueLen = 32 + +// stateTrackingStream is an implementation of quic.Stream that delegates +// to an underlying stream +// it takes care of proxying send and receive errors onto an implementation of +// the errorSetter interface (intended to be occupied by a datagrammer) +// it is also responsible for clearing the stream based on its ID from its +// parent connection, this is done through the streamClearer interface when +// both the send and receive sides are closed +type stateTrackingStream struct { + *quic.Stream + + sendDatagram func([]byte) error + hasData chan struct{} + queue [][]byte // TODO: use a ring buffer + + mx sync.Mutex + sendErr error + recvErr error + + clearer streamClearer +} + +var _ datagramStream = &stateTrackingStream{} + +type streamClearer interface { + clearStream(quic.StreamID) +} + +func newStateTrackingStream(s *quic.Stream, clearer streamClearer, sendDatagram func([]byte) error) *stateTrackingStream { + t := &stateTrackingStream{ + Stream: s, + clearer: clearer, + sendDatagram: sendDatagram, + hasData: make(chan struct{}, 1), + } + + context.AfterFunc(s.Context(), func() { + t.closeSend(context.Cause(s.Context())) + }) + + return t +} + +func (s *stateTrackingStream) closeSend(e error) { + s.mx.Lock() + defer s.mx.Unlock() + + // clear the stream the first time both the send + // and receive are finished + if s.sendErr == nil { + if s.recvErr != nil { + s.clearer.clearStream(s.StreamID()) + } + s.sendErr = e + } +} + +func (s *stateTrackingStream) closeReceive(e error) { + s.mx.Lock() + defer s.mx.Unlock() + + // clear the stream the first time both the send + // and receive are finished + if s.recvErr == nil { + if s.sendErr != nil { + s.clearer.clearStream(s.StreamID()) + } + s.recvErr = e + s.signalHasDatagram() + } +} + +func (s *stateTrackingStream) Close() error { + s.closeSend(errors.New("write on closed stream")) + return s.Stream.Close() +} + +func (s *stateTrackingStream) CancelWrite(e quic.StreamErrorCode) { + s.closeSend(&quic.StreamError{StreamID: s.StreamID(), ErrorCode: e}) + s.Stream.CancelWrite(e) +} + +func (s *stateTrackingStream) Write(b []byte) (int, error) { + n, err := s.Stream.Write(b) + if err != nil && !errors.Is(err, os.ErrDeadlineExceeded) { + s.closeSend(err) + } + return n, err +} + +func (s *stateTrackingStream) CancelRead(e quic.StreamErrorCode) { + s.closeReceive(&quic.StreamError{StreamID: s.StreamID(), ErrorCode: e}) + s.Stream.CancelRead(e) +} + +func (s *stateTrackingStream) Read(b []byte) (int, error) { + n, err := s.Stream.Read(b) + if err != nil && !errors.Is(err, os.ErrDeadlineExceeded) { + s.closeReceive(err) + } + return n, err +} + +func (s *stateTrackingStream) SendDatagram(b []byte) error { + s.mx.Lock() + sendErr := s.sendErr + s.mx.Unlock() + if sendErr != nil { + return sendErr + } + + return s.sendDatagram(b) +} + +func (s *stateTrackingStream) signalHasDatagram() { + select { + case s.hasData <- struct{}{}: + default: + } +} + +func (s *stateTrackingStream) enqueueDatagram(data []byte) { + s.mx.Lock() + defer s.mx.Unlock() + + if s.recvErr != nil { + return + } + if len(s.queue) >= streamDatagramQueueLen { + return + } + s.queue = append(s.queue, data) + s.signalHasDatagram() +} + +func (s *stateTrackingStream) ReceiveDatagram(ctx context.Context) ([]byte, error) { +start: + s.mx.Lock() + if len(s.queue) > 0 { + data := s.queue[0] + s.queue = s.queue[1:] + s.mx.Unlock() + return data, nil + } + if receiveErr := s.recvErr; receiveErr != nil { + s.mx.Unlock() + return nil, receiveErr + } + s.mx.Unlock() + + select { + case <-ctx.Done(): + return nil, context.Cause(ctx) + case <-s.hasData: + } + goto start +} + +func (s *stateTrackingStream) QUICStream() *quic.Stream { + return s.Stream +} diff --git a/third_party/quic-go/http3/state_tracking_stream_test.go b/third_party/quic-go/http3/state_tracking_stream_test.go new file mode 100644 index 0000000..f7d6f62 --- /dev/null +++ b/third_party/quic-go/http3/state_tracking_stream_test.go @@ -0,0 +1,319 @@ +package http3 + +import ( + "context" + "io" + "net" + "os" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func newStreamPair(t *testing.T) (client, server *quic.Stream) { + t.Helper() + + clientConn, serverConn := newConnPair(t) + serverStr, err := serverConn.OpenStream() + require.NoError(t, err) + // need to send something to the client to make it accept the stream + _, err = serverStr.Write([]byte{0}) + require.NoError(t, err) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientStr, err := clientConn.AcceptStream(ctx) + require.NoError(t, err) + clientStr.SetReadDeadline(time.Now().Add(time.Second)) + _, err = clientStr.Read([]byte{0}) + require.NoError(t, err) + clientStr.SetWriteDeadline(time.Time{}) + return clientStr, serverStr +} + +func canceledCtx() context.Context { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + return ctx +} + +func checkDatagramReceive(t *testing.T, str *stateTrackingStream) { + t.Helper() + _, err := str.ReceiveDatagram(canceledCtx()) + require.ErrorIs(t, err, context.Canceled) +} + +func checkDatagramSend(t *testing.T, str *stateTrackingStream) { + t.Helper() + require.NoError(t, str.SendDatagram([]byte("test"))) +} + +type mockStreamClearer struct { + cleared *quic.StreamID +} + +func (s *mockStreamClearer) clearStream(id quic.StreamID) { + s.cleared = &id +} + +func TestStateTrackingStreamRead(t *testing.T) { + t.Run("io.EOF", func(t *testing.T) { + testStateTrackingStreamRead(t, false) + }) + t.Run("remote stream reset", func(t *testing.T) { + testStateTrackingStreamRead(t, true) + }) +} + +func testStateTrackingStreamRead(t *testing.T, reset bool) { + client, server := newStreamPair(t) + + var clearer mockStreamClearer + str := newStateTrackingStream(client, &clearer, func(b []byte) error { return nil }) + + // deadline errors are ignored + client.SetReadDeadline(time.Now()) + _, err := str.Read(make([]byte, 3)) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Nil(t, clearer.cleared) + client.SetReadDeadline(time.Time{}) + + _, err = server.Write([]byte("foobar")) + require.NoError(t, err) + + if !reset { + server.Close() + + for range 3 { + _, err := str.Read([]byte{0}) + require.NoError(t, err) + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + } + } else { + server.CancelWrite(42) + } + + var expectedErr error + _, err = io.ReadAll(str) + if !reset { + require.NoError(t, err) + expectedErr = io.EOF + } else { + expectedErr = &quic.StreamError{Remote: true, StreamID: server.StreamID(), ErrorCode: 42} + require.ErrorIs(t, err, expectedErr) + } + require.Nil(t, clearer.cleared) + // the receive side registered the error + _, err = str.ReceiveDatagram(canceledCtx()) + require.ErrorIs(t, err, expectedErr) + // the send side is still open + require.NoError(t, str.SendDatagram([]byte("foo"))) +} + +func TestStateTrackingStreamRemoteCancelation(t *testing.T) { + client, server := newStreamPair(t) + + var clearer mockStreamClearer + str := newStateTrackingStream(client, &clearer, func(b []byte) error { return nil }) + + _, err := str.Write([]byte("foo")) + require.NoError(t, err) + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + checkDatagramSend(t, str) + + // deadline errors are ignored + client.SetWriteDeadline(time.Now()) + _, err = str.Write([]byte("baz")) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + checkDatagramSend(t, str) + client.SetWriteDeadline(time.Time{}) + + server.CancelRead(123) + + var writeErr error + require.Eventually(t, func() bool { + _, writeErr = str.Write([]byte("bar")) + return writeErr != nil + }, time.Second, scaleDuration(time.Millisecond)) + expectedErr := &quic.StreamError{Remote: true, StreamID: server.StreamID(), ErrorCode: 123} + require.ErrorIs(t, writeErr, expectedErr) + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + require.ErrorIs(t, str.SendDatagram([]byte("test")), expectedErr) +} + +func TestStateTrackingStreamLocalCancelation(t *testing.T) { + client, _ := newStreamPair(t) + + var clearer mockStreamClearer + str := newStateTrackingStream(client, &clearer, func(b []byte) error { return nil }) + + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + checkDatagramSend(t, str) + + str.CancelWrite(1337) + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + require.ErrorIs(t, str.SendDatagram([]byte("test")), &quic.StreamError{StreamID: client.StreamID(), ErrorCode: 1337}) +} + +func TestStateTrackingStreamClose(t *testing.T) { + client, _ := newStreamPair(t) + + var clearer mockStreamClearer + str := newStateTrackingStream(client, &clearer, func(b []byte) error { return nil }) + + require.Nil(t, clearer.cleared) + checkDatagramReceive(t, str) + checkDatagramSend(t, str) + + require.NoError(t, client.Close()) + require.Eventually(t, func() bool { + err := str.SendDatagram([]byte("test")) + if err == nil { + return false + } + require.ErrorIs(t, err, context.Canceled) + return true + }, time.Second, scaleDuration(5*time.Millisecond)) + + checkDatagramReceive(t, str) + require.Nil(t, clearer.cleared) +} + +func TestStateTrackingStreamReceiveThenSend(t *testing.T) { + client, server := newStreamPair(t) + + var clearer mockStreamClearer + str := newStateTrackingStream(client, &clearer, func(b []byte) error { return nil }) + + _, err := server.Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, server.Close()) + + _, err = io.ReadAll(str) + require.NoError(t, err) + + require.Nil(t, clearer.cleared) + _, err = str.ReceiveDatagram(canceledCtx()) + require.ErrorIs(t, err, io.EOF) + + client.CancelWrite(123) + + id := client.StreamID() + _, err = str.Write([]byte("bar")) + require.ErrorIs(t, err, &quic.StreamError{StreamID: id, ErrorCode: 123}) + require.ErrorIs(t, str.SendDatagram([]byte("test")), &quic.StreamError{StreamID: id, ErrorCode: 123}) + + require.Equal(t, &id, clearer.cleared) +} + +func TestStateTrackingStreamSendThenReceive(t *testing.T) { + client, server := newStreamPair(t) + + var clearer mockStreamClearer + str := newStateTrackingStream(client, &clearer, func(b []byte) error { return nil }) + + server.CancelRead(1234) + + var writeErr error + require.Eventually(t, func() bool { + _, writeErr = str.Write([]byte("bar")) + return writeErr != nil + }, time.Second, scaleDuration(time.Millisecond)) + id := server.StreamID() + expectedErr := &quic.StreamError{Remote: true, StreamID: id, ErrorCode: 1234} + require.ErrorIs(t, writeErr, expectedErr) + require.Nil(t, clearer.cleared) + require.ErrorIs(t, str.SendDatagram([]byte("test")), expectedErr) + + _, err := server.Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, server.Close()) + + _, err = io.ReadAll(str) + require.NoError(t, err) + _, err = str.ReceiveDatagram(canceledCtx()) + require.ErrorIs(t, err, io.EOF) + + require.Equal(t, &id, clearer.cleared) +} + +func TestDatagramReceiving(t *testing.T) { + client, _ := newStreamPair(t) + + str := newStateTrackingStream(client, nil, func(b []byte) error { return nil }) + type result struct { + data []byte + err error + } + + // Receive blocks until a datagram is received + resultChan := make(chan result) + go func() { + defer close(resultChan) + data, err := str.ReceiveDatagram(context.Background()) + resultChan <- result{data: data, err: err} + }() + + select { + case <-time.After(scaleDuration(10 * time.Millisecond)): + case <-resultChan: + t.Fatal("should not have received a datagram") + } + str.enqueueDatagram([]byte("foobar")) + + select { + case res := <-resultChan: + require.NoError(t, res.err) + require.Equal(t, []byte("foobar"), res.data) + case <-time.After(time.Second): + t.Fatal("should have received a datagram") + } + + // up to 32 datagrams can be queued + for i := range streamDatagramQueueLen + 1 { + str.enqueueDatagram([]byte{uint8(i)}) + } + for i := range streamDatagramQueueLen { + data, err := str.ReceiveDatagram(context.Background()) + require.NoError(t, err) + require.Equal(t, []byte{uint8(i)}, data) + } + + // Receive respects the context + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, err := str.ReceiveDatagram(ctx) + require.ErrorIs(t, err, context.Canceled) +} + +func TestDatagramSending(t *testing.T) { + var sendQueue [][]byte + errors := []error{nil, nil, assert.AnError} + client, _ := newStreamPair(t) + + str := newStateTrackingStream(client, nil, func(b []byte) error { + sendQueue = append(sendQueue, b) + err := errors[0] + errors = errors[1:] + return err + }) + require.NoError(t, str.SendDatagram([]byte("foo"))) + require.NoError(t, str.SendDatagram([]byte("bar"))) + require.ErrorIs(t, str.SendDatagram([]byte("baz")), assert.AnError) + require.Equal(t, [][]byte{[]byte("foo"), []byte("bar"), []byte("baz")}, sendQueue) + + str.closeSend(net.ErrClosed) + require.ErrorIs(t, str.SendDatagram([]byte("foobar")), net.ErrClosed) +} diff --git a/third_party/quic-go/http3/stream.go b/third_party/quic-go/http3/stream.go new file mode 100644 index 0000000..a1098b6 --- /dev/null +++ b/third_party/quic-go/http3/stream.go @@ -0,0 +1,406 @@ +package http3 + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptrace" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/quic-go/qpack" +) + +type datagramStream interface { + io.ReadWriteCloser + CancelRead(quic.StreamErrorCode) + CancelWrite(quic.StreamErrorCode) + StreamID() quic.StreamID + Context() context.Context + SetDeadline(time.Time) error + SetReadDeadline(time.Time) error + SetWriteDeadline(time.Time) error + SendDatagram(b []byte) error + ReceiveDatagram(ctx context.Context) ([]byte, error) + + QUICStream() *quic.Stream +} + +// A Stream is an HTTP/3 stream. +// +// When writing to and reading from the stream, data is framed in HTTP/3 DATA frames. +type Stream struct { + datagramStream + conn *rawConn + frameParser *frameParser + + buf []byte // used as a temporary buffer when writing the HTTP/3 frame headers + + bytesRemainingInFrame uint64 + + qlogger qlogwriter.Recorder + + parseTrailer func(io.Reader, *headersFrame) error + parsedTrailer bool +} + +func newStream( + str datagramStream, + conn *rawConn, + trace *httptrace.ClientTrace, + parseTrailer func(io.Reader, *headersFrame) error, + qlogger qlogwriter.Recorder, +) *Stream { + return &Stream{ + datagramStream: str, + conn: conn, + buf: make([]byte, 16), + qlogger: qlogger, + parseTrailer: parseTrailer, + frameParser: &frameParser{ + r: &tracingReader{Reader: str, trace: trace}, + streamID: str.StreamID(), + closeConn: conn.CloseWithError, + }, + } +} + +func (s *Stream) Read(b []byte) (int, error) { + if s.bytesRemainingInFrame == 0 { + parseLoop: + for { + frame, err := s.frameParser.ParseNext(s.qlogger) + if err != nil { + return 0, err + } + switch f := frame.(type) { + case *dataFrame: + if s.parsedTrailer { + return 0, errors.New("DATA frame received after trailers") + } + s.bytesRemainingInFrame = f.Length + break parseLoop + case *headersFrame: + if s.parsedTrailer { + maybeQlogInvalidHeadersFrame(s.qlogger, s.StreamID(), f.Length) + return 0, errors.New("additional HEADERS frame received after trailers") + } + s.parsedTrailer = true + return 0, s.parseTrailer(s.datagramStream, f) + default: + s.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "") + // parseNextFrame skips over unknown frame types + // Therefore, this condition is only entered when we parsed another known frame type. + return 0, fmt.Errorf("peer sent an unexpected frame: %T", f) + } + } + } + + var n int + var err error + if s.bytesRemainingInFrame < uint64(len(b)) { + n, err = s.datagramStream.Read(b[:s.bytesRemainingInFrame]) + } else { + n, err = s.datagramStream.Read(b) + } + s.bytesRemainingInFrame -= uint64(n) + return n, err +} + +func (s *Stream) hasMoreData() bool { + return s.bytesRemainingInFrame > 0 +} + +func (s *Stream) Write(b []byte) (int, error) { + s.buf = s.buf[:0] + s.buf = (&dataFrame{Length: uint64(len(b))}).Append(s.buf) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.FrameCreated{ + StreamID: s.StreamID(), + Raw: qlog.RawInfo{ + Length: len(s.buf) + len(b), + PayloadLength: len(b), + }, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }) + } + if _, err := s.datagramStream.Write(s.buf); err != nil { + return 0, err + } + return s.datagramStream.Write(b) +} + +func (s *Stream) writeUnframed(b []byte) (int, error) { + return s.datagramStream.Write(b) +} + +func (s *Stream) StreamID() quic.StreamID { + return s.datagramStream.StreamID() +} + +func (s *Stream) SendDatagram(b []byte) error { + // TODO: reject if datagrams are not negotiated (yet) + return s.datagramStream.SendDatagram(b) +} + +func (s *Stream) ReceiveDatagram(ctx context.Context) ([]byte, error) { + // TODO: reject if datagrams are not negotiated (yet) + return s.datagramStream.ReceiveDatagram(ctx) +} + +// A RequestStream is a low-level abstraction representing an HTTP/3 request stream. +// It decouples sending of the HTTP request from reading the HTTP response, allowing +// the application to optimistically use the stream (and, for example, send datagrams) +// before receiving the response. +// +// This is only needed for advanced use case, e.g. WebTransport and the various +// MASQUE proxying protocols. +type RequestStream struct { + str *Stream + + responseBody io.ReadCloser // set by ReadResponse + + decoder *qpack.Decoder + requestWriter *requestWriter + maxHeaderBytes int + reqDone chan<- struct{} + disableCompression bool + response *http.Response + + sentRequest bool + requestedGzip bool + isConnect bool +} + +func newRequestStream( + str *Stream, + requestWriter *requestWriter, + reqDone chan<- struct{}, + decoder *qpack.Decoder, + disableCompression bool, + maxHeaderBytes int, + rsp *http.Response, +) *RequestStream { + return &RequestStream{ + str: str, + requestWriter: requestWriter, + reqDone: reqDone, + decoder: decoder, + disableCompression: disableCompression, + maxHeaderBytes: maxHeaderBytes, + response: rsp, + } +} + +// Read reads data from the underlying stream. +// +// It can only be used after the request has been sent (using SendRequestHeader) +// and the response has been consumed (using ReadResponse). +func (s *RequestStream) Read(b []byte) (int, error) { + if s.responseBody == nil { + return 0, errors.New("http3: invalid use of RequestStream.Read before ReadResponse") + } + return s.responseBody.Read(b) +} + +// StreamID returns the QUIC stream ID of the underlying QUIC stream. +func (s *RequestStream) StreamID() quic.StreamID { + return s.str.StreamID() +} + +// Write writes data to the stream. +// +// It can only be used after the request has been sent (using SendRequestHeader). +func (s *RequestStream) Write(b []byte) (int, error) { + if !s.sentRequest { + return 0, errors.New("http3: invalid use of RequestStream.Write before SendRequestHeader") + } + return s.str.Write(b) +} + +// Close closes the send-direction of the stream. +// It does not close the receive-direction of the stream. +func (s *RequestStream) Close() error { + return s.str.Close() +} + +// CancelRead aborts receiving on this stream. +// See [quic.Stream.CancelRead] for more details. +func (s *RequestStream) CancelRead(errorCode quic.StreamErrorCode) { + s.str.CancelRead(errorCode) +} + +// CancelWrite aborts sending on this stream. +// See [quic.Stream.CancelWrite] for more details. +func (s *RequestStream) CancelWrite(errorCode quic.StreamErrorCode) { + s.str.CancelWrite(errorCode) +} + +// Context returns a context derived from the underlying QUIC stream's context. +// See [quic.Stream.Context] for more details. +func (s *RequestStream) Context() context.Context { + return s.str.Context() +} + +// SetReadDeadline sets the deadline for Read calls. +func (s *RequestStream) SetReadDeadline(t time.Time) error { + return s.str.SetReadDeadline(t) +} + +// SetWriteDeadline sets the deadline for Write calls. +func (s *RequestStream) SetWriteDeadline(t time.Time) error { + return s.str.SetWriteDeadline(t) +} + +// SetDeadline sets the read and write deadlines associated with the stream. +// It is equivalent to calling both SetReadDeadline and SetWriteDeadline. +func (s *RequestStream) SetDeadline(t time.Time) error { + return s.str.SetDeadline(t) +} + +// SendDatagrams send a new HTTP Datagram (RFC 9297). +// +// It is only possible to send datagrams if the server enabled support for this extension. +// It is recommended (though not required) to send the request before calling this method, +// as the server might drop datagrams which it can't associate with an existing request. +func (s *RequestStream) SendDatagram(b []byte) error { + return s.str.SendDatagram(b) +} + +// ReceiveDatagram receives HTTP Datagrams (RFC 9297). +// +// It is only possible if support for HTTP Datagrams was enabled, using the EnableDatagram +// option on the [Transport]. +func (s *RequestStream) ReceiveDatagram(ctx context.Context) ([]byte, error) { + return s.str.ReceiveDatagram(ctx) +} + +// SendRequestHeader sends the HTTP request. +// +// It can only used for requests that don't have a request body. +// It is invalid to call it more than once. +// It is invalid to call it after Write has been called. +func (s *RequestStream) SendRequestHeader(req *http.Request) error { + if req.Body != nil && req.Body != http.NoBody { + return errors.New("http3: invalid use of RequestStream.SendRequestHeader with a request that has a request body") + } + return s.sendRequestHeader(req) +} + +func (s *RequestStream) sendRequestHeader(req *http.Request) error { + if s.sentRequest { + return errors.New("http3: invalid duplicate use of RequestStream.SendRequestHeader") + } + if !s.disableCompression && req.Method != http.MethodHead && + req.Header.Get("Accept-Encoding") == "" && req.Header.Get("Range") == "" { + s.requestedGzip = true + } + s.isConnect = req.Method == http.MethodConnect + s.sentRequest = true + return s.requestWriter.WriteRequestHeader(s.str.datagramStream, req, s.requestedGzip, s.str.StreamID(), s.str.qlogger) +} + +// sendRequestTrailer sends request trailers to the stream. +// It should be called after the request body has been fully written. +func (s *RequestStream) sendRequestTrailer(req *http.Request) error { + return s.requestWriter.WriteRequestTrailer(s.str.datagramStream, req, s.str.StreamID(), s.str.qlogger) +} + +// ReadResponse reads the HTTP response from the stream. +// +// It must be called after sending the request (using SendRequestHeader). +// It is invalid to call it more than once. +// It doesn't set Response.Request and Response.TLS. +// It is invalid to call it after Read has been called. +func (s *RequestStream) ReadResponse() (*http.Response, error) { + if !s.sentRequest { + return nil, errors.New("http3: invalid use of RequestStream.ReadResponse before SendRequestHeader") + } + frame, err := s.str.frameParser.ParseNext(s.str.qlogger) + if err != nil { + s.str.CancelRead(quic.StreamErrorCode(ErrCodeFrameError)) + s.str.CancelWrite(quic.StreamErrorCode(ErrCodeFrameError)) + return nil, fmt.Errorf("http3: parsing frame failed: %w", err) + } + hf, ok := frame.(*headersFrame) + if !ok { + s.str.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "expected first frame to be a HEADERS frame") + return nil, errors.New("http3: expected first frame to be a HEADERS frame") + } + if hf.Length > uint64(s.maxHeaderBytes) { + maybeQlogInvalidHeadersFrame(s.str.qlogger, s.str.StreamID(), hf.Length) + s.str.CancelRead(quic.StreamErrorCode(ErrCodeFrameError)) + s.str.CancelWrite(quic.StreamErrorCode(ErrCodeFrameError)) + return nil, fmt.Errorf("http3: HEADERS frame too large: %d bytes (max: %d)", hf.Length, s.maxHeaderBytes) + } + headerBlock := make([]byte, hf.Length) + if _, err := io.ReadFull(s.str.datagramStream, headerBlock); err != nil { + maybeQlogInvalidHeadersFrame(s.str.qlogger, s.str.StreamID(), hf.Length) + s.str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + s.str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete)) + return nil, fmt.Errorf("http3: failed to read response headers: %w", err) + } + decodeFn := s.decoder.Decode(headerBlock) + var hfs []qpack.HeaderField + if s.str.qlogger != nil { + hfs = make([]qpack.HeaderField, 0, 16) + } + res := s.response + err = updateResponseFromHeaders(res, decodeFn, s.maxHeaderBytes, &hfs) + if s.str.qlogger != nil { + qlogParsedHeadersFrame(s.str.qlogger, s.str.StreamID(), hf, hfs) + } + if err != nil { + errCode := ErrCodeMessageError + var qpackErr *qpackError + if errors.As(err, &qpackErr) { + errCode = ErrCodeQPACKDecompressionFailed + } + s.str.CancelRead(quic.StreamErrorCode(errCode)) + s.str.CancelWrite(quic.StreamErrorCode(errCode)) + return nil, fmt.Errorf("http3: invalid response: %w", err) + } + + // Check that the server doesn't send more data in DATA frames than indicated by the Content-Length header (if set). + // See section 4.1.2 of RFC 9114. + respBody := newResponseBody(s.str, res.ContentLength, s.reqDone) + + // Rules for when to set Content-Length are defined in https://tools.ietf.org/html/rfc7230#section-3.3.2. + isInformational := res.StatusCode >= 100 && res.StatusCode < 200 + isNoContent := res.StatusCode == http.StatusNoContent + isSuccessfulConnect := s.isConnect && res.StatusCode >= 200 && res.StatusCode < 300 + if (isInformational || isNoContent || isSuccessfulConnect) && res.ContentLength == -1 { + res.ContentLength = 0 + } + if s.requestedGzip && res.Header.Get("Content-Encoding") == "gzip" { + res.Header.Del("Content-Encoding") + res.Header.Del("Content-Length") + res.ContentLength = -1 + s.responseBody = newGzipReader(respBody) + res.Uncompressed = true + } else { + s.responseBody = respBody + } + res.Body = s.responseBody + return res, nil +} + +type tracingReader struct { + io.Reader + readFirst bool + trace *httptrace.ClientTrace +} + +func (r *tracingReader) Read(b []byte) (int, error) { + n, err := r.Reader.Read(b) + if n > 0 && !r.readFirst { + traceGotFirstResponseByte(r.trace) + r.readFirst = true + } + return n, err +} diff --git a/third_party/quic-go/http3/stream_test.go b/third_party/quic-go/http3/stream_test.go new file mode 100644 index 0000000..c7bf37f --- /dev/null +++ b/third_party/quic-go/http3/stream_test.go @@ -0,0 +1,235 @@ +package http3 + +import ( + "bytes" + "context" + "io" + "math" + "net/http" + "net/http/httptest" + "net/http/httptrace" + "strings" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/quic-go/qpack" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func getDataFrame(data []byte) []byte { + b := (&dataFrame{Length: uint64(len(data))}).Append(nil) + return append(b, data...) +} + +func TestStreamReadDataFrames(t *testing.T) { + var buf bytes.Buffer + mockCtrl := gomock.NewController(t) + qstr := NewMockDatagramStream(mockCtrl) + qstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + qstr.EXPECT().Write(gomock.Any()).DoAndReturn(buf.Write).AnyTimes() + qstr.EXPECT().Read(gomock.Any()).DoAndReturn(buf.Read).AnyTimes() + + var eventRecorder events.Recorder + clientConn, _ := newConnPair(t, withClientRecorder(&eventRecorder)) + str := newStream( + qstr, + newRawConn(clientConn, false, nil, nil, &eventRecorder, nil), + nil, + func(io.Reader, *headersFrame) error { return nil }, + &eventRecorder, + ) + + buf.Write(getDataFrame([]byte("foobar"))) + b := make([]byte, 3) + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("foo"), b) + n, err = str.Read(b) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("bar"), b) + + expectedLen, _ := expectedFrameLength(t, &dataFrame{Length: 6}) + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameParsed{ + StreamID: 42, + Raw: qlog.RawInfo{Length: expectedLen, PayloadLength: 6}, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }, + }, + eventRecorder.Events(qlog.FrameParsed{}), + ) + eventRecorder.Clear() + + buf.Write(getDataFrame([]byte("baz"))) + b = make([]byte, 10) + n, err = str.Read(b) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("baz"), b[:n]) + require.Len(t, eventRecorder.Events(qlog.FrameParsed{}), 1) + eventRecorder.Clear() + + buf.Write(getDataFrame([]byte("lorem"))) + buf.Write(getDataFrame([]byte("ipsum"))) + + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, "loremipsum", string(data)) + require.Len(t, eventRecorder.Events(qlog.FrameParsed{}), 2) + eventRecorder.Clear() + + // invalid frame + buf.Write([]byte("invalid")) + _, err = str.Read([]byte{0}) + require.Error(t, err) +} + +func TestStreamInvalidFrame(t *testing.T) { + var buf bytes.Buffer + b := (&settingsFrame{}).Append(nil) + buf.Write(b) + + mockCtrl := gomock.NewController(t) + qstr := NewMockDatagramStream(mockCtrl) + qstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + qstr.EXPECT().Write(gomock.Any()).DoAndReturn(buf.Write).AnyTimes() + qstr.EXPECT().Read(gomock.Any()).DoAndReturn(buf.Read).AnyTimes() + clientConn, serverConn := newConnPair(t) + + str := newStream( + qstr, + newRawConn(clientConn, false, nil, nil, nil, nil), + nil, + func(io.Reader, *headersFrame) error { return nil }, + nil, + ) + + _, err := str.Read([]byte{0}) + require.ErrorContains(t, err, "peer sent an unexpected frame") + + select { + case <-serverConn.Context().Done(): + var appErr *quic.ApplicationError + require.ErrorAs(t, context.Cause(serverConn.Context()), &appErr) + require.Equal(t, quic.ApplicationErrorCode(ErrCodeFrameUnexpected), appErr.ErrorCode) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestStreamWrite(t *testing.T) { + var buf bytes.Buffer + mockCtrl := gomock.NewController(t) + qstr := NewMockDatagramStream(mockCtrl) + qstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + qstr.EXPECT().Write(gomock.Any()).DoAndReturn(buf.Write).AnyTimes() + var eventRecorder events.Recorder + str := newStream(qstr, nil, nil, func(io.Reader, *headersFrame) error { return nil }, &eventRecorder) + str.Write([]byte("foo")) + str.Write([]byte("foobar")) + + fp := frameParser{r: &buf} + f, err := fp.ParseNext(nil) + require.NoError(t, err) + f1Len, f1PayloadLen := expectedFrameLength(t, &dataFrame{Length: 3}) + require.Equal(t, &dataFrame{Length: 3}, f) + b := make([]byte, 3) + _, err = io.ReadFull(&buf, b) + require.NoError(t, err) + require.Equal(t, []byte("foo"), b) + + fp = frameParser{r: &buf} + f, err = fp.ParseNext(nil) + require.NoError(t, err) + f2Len, f2PayloadLen := expectedFrameLength(t, &dataFrame{Length: 6}) + require.Equal(t, &dataFrame{Length: 6}, f) + b = make([]byte, 6) + _, err = io.ReadFull(&buf, b) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), b) + + require.Equal(t, + []qlogwriter.Event{ + qlog.FrameCreated{ + StreamID: 42, + Raw: qlog.RawInfo{Length: f1Len, PayloadLength: f1PayloadLen}, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }, + qlog.FrameCreated{ + StreamID: 42, + Raw: qlog.RawInfo{Length: f2Len, PayloadLength: f2PayloadLen}, + Frame: qlog.Frame{Frame: qlog.DataFrame{}}, + }, + }, + eventRecorder.Events(qlog.FrameCreated{}), + ) +} + +func TestRequestStream(t *testing.T) { + mockCtrl := gomock.NewController(t) + qstr := NewMockDatagramStream(mockCtrl) + qstr.EXPECT().StreamID().Return(quic.StreamID(42)).AnyTimes() + requestWriter := newRequestWriter() + clientConn, _ := newConnPair(t) + str := newRequestStream( + newStream( + qstr, + newRawConn(clientConn, false, nil, nil, nil, nil), + &httptrace.ClientTrace{}, + func(io.Reader, *headersFrame) error { return nil }, + nil, + ), + requestWriter, + make(chan struct{}), + qpack.NewDecoder(), + true, + math.MaxInt, + &http.Response{}, + ) + + _, err := str.Read([]byte{0}) + require.EqualError(t, err, "http3: invalid use of RequestStream.Read before ReadResponse") + _, err = str.Write([]byte{0}) + require.EqualError(t, err, "http3: invalid use of RequestStream.Write before SendRequestHeader") + + // calling ReadResponse before SendRequestHeader is not valid + _, err = str.ReadResponse() + require.EqualError(t, err, "http3: invalid use of RequestStream.ReadResponse before SendRequestHeader") + // SendRequestHeader can't be used for requests that have a request body + require.EqualError(t, + str.SendRequestHeader( + httptest.NewRequest(http.MethodGet, "https://quic-go.net", strings.NewReader("foobar")), + ), + "http3: invalid use of RequestStream.SendRequestHeader with a request that has a request body", + ) + + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net", nil) + qstr.EXPECT().Write(gomock.Any()).AnyTimes() + require.NoError(t, str.SendRequestHeader(req)) + // duplicate calls are not allowed + require.EqualError(t, str.SendRequestHeader(req), "http3: invalid duplicate use of RequestStream.SendRequestHeader") + + buf := bytes.NewBuffer(encodeResponse(t, http.StatusOK)) + buf.Write((&dataFrame{Length: 6}).Append(nil)) + buf.Write([]byte("foobar")) + qstr.EXPECT().Read(gomock.Any()).DoAndReturn(buf.Read).AnyTimes() + rsp, err := str.ReadResponse() + require.NoError(t, err) + require.Equal(t, http.StatusOK, rsp.StatusCode) + + b := make([]byte, 10) + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 6, n) + require.Equal(t, []byte("foobar"), b[:n]) +} diff --git a/third_party/quic-go/http3/trace.go b/third_party/quic-go/http3/trace.go new file mode 100644 index 0000000..78962e9 --- /dev/null +++ b/third_party/quic-go/http3/trace.go @@ -0,0 +1,105 @@ +package http3 + +import ( + "crypto/tls" + "net" + "net/http/httptrace" + "net/textproto" + "time" + + "github.com/apernet/quic-go" +) + +func traceGetConn(trace *httptrace.ClientTrace, hostPort string) { + if trace != nil && trace.GetConn != nil { + trace.GetConn(hostPort) + } +} + +// fakeConn is a wrapper for quic.EarlyConnection +// because the quic connection does not implement net.Conn. +type fakeConn struct { + conn *quic.Conn +} + +func (c *fakeConn) Close() error { panic("connection operation prohibited") } +func (c *fakeConn) Read(p []byte) (int, error) { panic("connection operation prohibited") } +func (c *fakeConn) Write(p []byte) (int, error) { panic("connection operation prohibited") } +func (c *fakeConn) SetDeadline(t time.Time) error { panic("connection operation prohibited") } +func (c *fakeConn) SetReadDeadline(t time.Time) error { panic("connection operation prohibited") } +func (c *fakeConn) SetWriteDeadline(t time.Time) error { panic("connection operation prohibited") } +func (c *fakeConn) RemoteAddr() net.Addr { return c.conn.RemoteAddr() } +func (c *fakeConn) LocalAddr() net.Addr { return c.conn.LocalAddr() } + +func traceGotConn(trace *httptrace.ClientTrace, conn *quic.Conn, reused bool) { + if trace != nil && trace.GotConn != nil { + trace.GotConn(httptrace.GotConnInfo{ + Conn: &fakeConn{conn: conn}, + Reused: reused, + }) + } +} + +func traceGotFirstResponseByte(trace *httptrace.ClientTrace) { + if trace != nil && trace.GotFirstResponseByte != nil { + trace.GotFirstResponseByte() + } +} + +func traceGot1xxResponse(trace *httptrace.ClientTrace, code int, header textproto.MIMEHeader) { + if trace != nil && trace.Got1xxResponse != nil { + trace.Got1xxResponse(code, header) + } +} + +func traceGot100Continue(trace *httptrace.ClientTrace) { + if trace != nil && trace.Got100Continue != nil { + trace.Got100Continue() + } +} + +func traceHasWroteHeaderField(trace *httptrace.ClientTrace) bool { + return trace != nil && trace.WroteHeaderField != nil +} + +func traceWroteHeaderField(trace *httptrace.ClientTrace, k, v string) { + if trace != nil && trace.WroteHeaderField != nil { + trace.WroteHeaderField(k, []string{v}) + } +} + +func traceWroteHeaders(trace *httptrace.ClientTrace) { + if trace != nil && trace.WroteHeaders != nil { + trace.WroteHeaders() + } +} + +func traceWroteRequest(trace *httptrace.ClientTrace, err error) { + if trace != nil && trace.WroteRequest != nil { + trace.WroteRequest(httptrace.WroteRequestInfo{Err: err}) + } +} + +func traceConnectStart(trace *httptrace.ClientTrace, network, addr string) { + if trace != nil && trace.ConnectStart != nil { + trace.ConnectStart(network, addr) + } +} + +func traceConnectDone(trace *httptrace.ClientTrace, network, addr string, err error) { + if trace != nil && trace.ConnectDone != nil { + trace.ConnectDone(network, addr, err) + } +} + +func traceTLSHandshakeStart(trace *httptrace.ClientTrace) { + if trace != nil && trace.TLSHandshakeStart != nil { + trace.TLSHandshakeStart() + } +} + +func traceTLSHandshakeDone(trace *httptrace.ClientTrace, state tls.ConnectionState, err error) { + if trace != nil && trace.TLSHandshakeDone != nil { + trace.TLSHandshakeDone(state, err) + } +} diff --git a/third_party/quic-go/http3/transport.go b/third_party/quic-go/http3/transport.go new file mode 100644 index 0000000..7d4f80a --- /dev/null +++ b/third_party/quic-go/http3/transport.go @@ -0,0 +1,538 @@ +package http3 + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "log/slog" + "net" + "net/http" + "net/http/httptrace" + "net/url" + "strings" + "sync" + "sync/atomic" + + "golang.org/x/net/http/httpguts" + + "github.com/apernet/quic-go" +) + +// Settings are HTTP/3 settings that apply to the underlying connection. +type Settings struct { + // Support for HTTP/3 datagrams (RFC 9297) + EnableDatagrams bool + // Extended CONNECT, RFC 9220 + EnableExtendedConnect bool + // Other settings, defined by the application + Other map[uint64]uint64 +} + +// RoundTripOpt are options for the Transport.RoundTripOpt method. +type RoundTripOpt struct { + // OnlyCachedConn controls whether the Transport may create a new QUIC connection. + // If set true and no cached connection is available, RoundTripOpt will return ErrNoCachedConn. + OnlyCachedConn bool +} + +type clientConn interface { + OpenRequestStream(context.Context) (*RequestStream, error) + RoundTrip(*http.Request) (*http.Response, error) + handleUnidirectionalStream(*quic.ReceiveStream) +} + +type roundTripperWithCount struct { + cancel context.CancelFunc + dialing chan struct{} // closed as soon as quic.Dial(Early) returned + dialErr error + conn *quic.Conn + clientConn clientConn + + useCount atomic.Int64 +} + +func (r *roundTripperWithCount) Close() error { + r.cancel() + <-r.dialing + if r.conn != nil { + return r.conn.CloseWithError(0, "") + } + return nil +} + +// Transport implements the http.RoundTripper interface +type Transport struct { + // TLSClientConfig specifies the TLS configuration to use with + // tls.Client. If nil, the default configuration is used. + TLSClientConfig *tls.Config + + // QUICConfig is the quic.Config used for dialing new connections. + // If nil, reasonable default values will be used. + QUICConfig *quic.Config + + // Dial specifies an optional dial function for creating QUIC + // connections for requests. + // If Dial is nil, a UDPConn will be created at the first request + // and will be reused for subsequent connections to other servers. + Dial func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) + + // Enable support for HTTP/3 datagrams (RFC 9297). + // If a QUICConfig is set, datagram support also needs to be enabled on the QUIC layer by setting EnableDatagrams. + EnableDatagrams bool + + // Additional HTTP/3 settings. + // It is invalid to specify any settings defined by RFC 9114 (HTTP/3) and RFC 9297 (HTTP Datagrams). + AdditionalSettings map[uint64]uint64 + + // MaxResponseHeaderBytes specifies a limit on how many response bytes are + // allowed in the server's response header. + // Zero means to use a default limit. + MaxResponseHeaderBytes int + + // DisableCompression, if true, prevents the Transport from requesting compression with an + // "Accept-Encoding: gzip" request header when the Request contains no existing Accept-Encoding value. + // If the Transport requests gzip on its own and gets a gzipped response, it's transparently + // decoded in the Response.Body. + // However, if the user explicitly requested gzip it is not automatically uncompressed. + DisableCompression bool + + Logger *slog.Logger + + mutex sync.Mutex + + initOnce sync.Once + initErr error + + newClientConn func(*quic.Conn) clientConn + + clients map[string]*roundTripperWithCount + transport *quic.Transport + closed bool +} + +var ( + _ http.RoundTripper = &Transport{} + _ io.Closer = &Transport{} +) + +var ( + // ErrNoCachedConn is returned when Transport.OnlyCachedConn is set + ErrNoCachedConn = errors.New("http3: no cached connection was available") + // ErrTransportClosed is returned when attempting to use a closed Transport + ErrTransportClosed = errors.New("http3: transport is closed") +) + +func (t *Transport) init() error { + if t.newClientConn == nil { + t.newClientConn = func(conn *quic.Conn) clientConn { + return newClientConn( + conn, + t.EnableDatagrams, + t.AdditionalSettings, + t.MaxResponseHeaderBytes, + t.DisableCompression, + t.Logger, + ) + } + } + if t.QUICConfig == nil { + t.QUICConfig = defaultQuicConfig.Clone() + t.QUICConfig.EnableDatagrams = t.EnableDatagrams + } + if t.EnableDatagrams && !t.QUICConfig.EnableDatagrams { + return errors.New("HTTP Datagrams enabled, but QUIC Datagrams disabled") + } + if len(t.QUICConfig.Versions) == 0 { + t.QUICConfig = t.QUICConfig.Clone() + t.QUICConfig.Versions = []quic.Version{quic.SupportedVersions()[0]} + } + if len(t.QUICConfig.Versions) != 1 { + return errors.New("can only use a single QUIC version for dialing a HTTP/3 connection") + } + if t.QUICConfig.MaxIncomingStreams == 0 { + t.QUICConfig.MaxIncomingStreams = -1 // don't allow any bidirectional streams + } + if t.Dial == nil { + udpConn, err := net.ListenUDP("udp", nil) + if err != nil { + return err + } + t.transport = &quic.Transport{Conn: udpConn} + } + return nil +} + +// RoundTripOpt is like RoundTrip, but takes options. +func (t *Transport) RoundTripOpt(req *http.Request, opt RoundTripOpt) (*http.Response, error) { + rsp, err := t.roundTripOpt(req, opt) + if err != nil { + if req.Body != nil { + req.Body.Close() + } + return nil, err + } + return rsp, nil +} + +func (t *Transport) roundTripOpt(req *http.Request, opt RoundTripOpt) (*http.Response, error) { + t.initOnce.Do(func() { t.initErr = t.init() }) + if t.initErr != nil { + return nil, t.initErr + } + + if req.URL == nil { + return nil, errors.New("http3: nil Request.URL") + } + if req.URL.Scheme != "https" { + return nil, fmt.Errorf("http3: unsupported protocol scheme: %s", req.URL.Scheme) + } + if req.URL.Host == "" { + return nil, errors.New("http3: no Host in request URL") + } + if req.Header == nil { + return nil, errors.New("http3: nil Request.Header") + } + if req.Method != "" && !validMethod(req.Method) { + return nil, fmt.Errorf("http3: invalid method %q", req.Method) + } + for k, vv := range req.Header { + if !httpguts.ValidHeaderFieldName(k) { + return nil, fmt.Errorf("http3: invalid http header field name %q", k) + } + for _, v := range vv { + if !httpguts.ValidHeaderFieldValue(v) { + return nil, fmt.Errorf("http3: invalid http header field value for key %q", k) + } + } + } + + return t.doRoundTripOpt(req, opt, false) +} + +func (t *Transport) doRoundTripOpt(req *http.Request, opt RoundTripOpt, isRetried bool) (*http.Response, error) { + hostname := authorityAddr(hostnameFromURL(req.URL)) + trace := httptrace.ContextClientTrace(req.Context()) + traceGetConn(trace, hostname) + cl, isReused, err := t.getClient(req.Context(), hostname, opt.OnlyCachedConn) + if err != nil { + return nil, err + } + + select { + case <-cl.dialing: + case <-req.Context().Done(): + return nil, context.Cause(req.Context()) + } + + if cl.dialErr != nil { + t.removeClient(hostname) + return nil, cl.dialErr + } + defer cl.useCount.Add(-1) + traceGotConn(trace, cl.conn, isReused) + rsp, err := cl.clientConn.RoundTrip(req) + if err != nil { + // request aborted due to context cancellation + select { + case <-req.Context().Done(): + return nil, err + default: + } + if isRetried { + return nil, err + } + + t.removeClient(hostname) + req, err = canRetryRequest(err, req) + if err != nil { + return nil, err + } + return t.doRoundTripOpt(req, opt, true) + } + return rsp, nil +} + +func canRetryRequest(err error, req *http.Request) (*http.Request, error) { + // error occurred while opening the stream, we can be sure that the request wasn't sent out + var connErr *errConnUnusable + if errors.As(err, &connErr) { + return req, nil + } + + // If the request stream is reset, we can only be sure that the request wasn't processed + // if the error code is H3_REQUEST_REJECTED. + var e *Error + if !errors.As(err, &e) || e.ErrorCode != ErrCodeRequestRejected { + return nil, err + } + // if the body is nil (or http.NoBody), it's safe to reuse this request and its body + if req.Body == nil || req.Body == http.NoBody { + return req, nil + } + // if the request body can be reset back to its original state via req.GetBody, do that + if req.GetBody != nil { + newBody, err := req.GetBody() + if err != nil { + return nil, err + } + reqCopy := *req + reqCopy.Body = newBody + req = &reqCopy + return &reqCopy, nil + } + return nil, fmt.Errorf("http3: Transport: cannot retry err [%w] after Request.Body was written; define Request.GetBody to avoid this error", err) +} + +// RoundTrip does a round trip. +func (t *Transport) RoundTrip(req *http.Request) (*http.Response, error) { + return t.RoundTripOpt(req, RoundTripOpt{}) +} + +func (t *Transport) getClient(ctx context.Context, hostname string, onlyCached bool) (rtc *roundTripperWithCount, isReused bool, err error) { + t.mutex.Lock() + defer t.mutex.Unlock() + if t.closed { + return nil, false, ErrTransportClosed + } + + if t.clients == nil { + t.clients = make(map[string]*roundTripperWithCount) + } + + cl, ok := t.clients[hostname] + if !ok { + if onlyCached { + return nil, false, ErrNoCachedConn + } + ctx, cancel := context.WithCancel(ctx) + cl = &roundTripperWithCount{ + dialing: make(chan struct{}), + cancel: cancel, + } + go func() { + defer close(cl.dialing) + defer cancel() + conn, rt, err := t.dial(ctx, hostname) + if err != nil { + cl.dialErr = err + return + } + cl.conn = conn + cl.clientConn = rt + }() + t.clients[hostname] = cl + } + select { + case <-cl.dialing: + if cl.dialErr != nil { + delete(t.clients, hostname) + return nil, false, cl.dialErr + } + select { + case <-cl.conn.HandshakeComplete(): + isReused = true + default: + } + default: + } + cl.useCount.Add(1) + return cl, isReused, nil +} + +func (t *Transport) dial(ctx context.Context, hostname string) (*quic.Conn, clientConn, error) { + var tlsConf *tls.Config + if t.TLSClientConfig == nil { + tlsConf = &tls.Config{} + } else { + tlsConf = t.TLSClientConfig.Clone() + } + if tlsConf.ServerName == "" { + sni, _, err := net.SplitHostPort(hostname) + if err != nil { + // It's ok if net.SplitHostPort returns an error - it could be a hostname/IP address without a port. + sni = hostname + } + tlsConf.ServerName = sni + } + // Replace existing ALPNs by H3 + tlsConf.NextProtos = []string{NextProtoH3} + + dial := t.Dial + if dial == nil { + dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + network := "udp" + udpAddr, err := t.resolveUDPAddr(ctx, network, addr) + if err != nil { + return nil, err + } + trace := httptrace.ContextClientTrace(ctx) + traceConnectStart(trace, network, udpAddr.String()) + traceTLSHandshakeStart(trace) + conn, err := t.transport.DialEarly(ctx, udpAddr, tlsCfg, cfg) + var state tls.ConnectionState + if conn != nil { + state = conn.ConnectionState().TLS + } + traceTLSHandshakeDone(trace, state, err) + traceConnectDone(trace, network, udpAddr.String(), err) + return conn, err + } + } + conn, err := dial(ctx, hostname, tlsConf, t.QUICConfig) + if err != nil { + return nil, nil, err + } + clientConn := t.newClientConn(conn) + go func() { + for { + str, err := conn.AcceptUniStream(context.Background()) + if err != nil { + return + } + go clientConn.handleUnidirectionalStream(str) + } + }() + return conn, clientConn, nil +} + +func (t *Transport) resolveUDPAddr(ctx context.Context, network, addr string) (*net.UDPAddr, error) { + host, portStr, err := net.SplitHostPort(addr) + if err != nil { + return nil, err + } + port, err := net.LookupPort(network, portStr) + if err != nil { + return nil, err + } + resolver := net.DefaultResolver + ipAddrs, err := resolver.LookupIPAddr(ctx, host) + if err != nil { + return nil, err + } + addrs := addrList(ipAddrs) + ip := addrs.forResolve(network, addr) + return &net.UDPAddr{IP: ip.IP, Port: port, Zone: ip.Zone}, nil +} + +func (t *Transport) removeClient(hostname string) { + t.mutex.Lock() + defer t.mutex.Unlock() + if t.clients == nil { + return + } + delete(t.clients, hostname) +} + +// NewClientConn creates a new HTTP/3 client connection on top of a QUIC connection. +// Most users should use RoundTrip instead of creating a connection directly. +// Specifically, it is not needed to perform GET, POST, HEAD and CONNECT requests. +// +// Obtaining a ClientConn is only needed for more advanced use cases, such as +// using Extended CONNECT for WebTransport or the various MASQUE protocols. +func (t *Transport) NewClientConn(conn *quic.Conn) *ClientConn { + c := newClientConn( + conn, + t.EnableDatagrams, + t.AdditionalSettings, + t.MaxResponseHeaderBytes, + t.DisableCompression, + t.Logger, + ) + go func() { + for { + str, err := conn.AcceptUniStream(context.Background()) + if err != nil { + return + } + go c.handleUnidirectionalStream(str) + } + }() + return c +} + +// NewRawClientConn creates a new low-level HTTP/3 client connection on top of a QUIC connection. +// Unlike NewClientConn, the returned RawClientConn allows the application to take control +// of the stream accept loops, by calling HandleUnidirectionalStream for incoming unidirectional +// streams and HandleBidirectionalStream for incoming bidirectional streams. +func (t *Transport) NewRawClientConn(conn *quic.Conn) *RawClientConn { + return &RawClientConn{ + ClientConn: newClientConn( + conn, + t.EnableDatagrams, + t.AdditionalSettings, + t.MaxResponseHeaderBytes, + t.DisableCompression, + t.Logger, + ), + } +} + +// Close closes the QUIC connections that this Transport has used. +// A Transport cannot be used after it has been closed. +func (t *Transport) Close() error { + t.mutex.Lock() + defer t.mutex.Unlock() + for _, cl := range t.clients { + if err := cl.Close(); err != nil { + return err + } + } + t.clients = nil + if t.transport != nil { + if err := t.transport.Close(); err != nil { + return err + } + if err := t.transport.Conn.Close(); err != nil { + return err + } + t.transport = nil + } + t.closed = true + return nil +} + +func hostnameFromURL(url *url.URL) string { + if url != nil { + return url.Host + } + return "" +} + +func validMethod(method string) bool { + /* + Method = "OPTIONS" ; Section 9.2 + | "GET" ; Section 9.3 + | "HEAD" ; Section 9.4 + | "POST" ; Section 9.5 + | "PUT" ; Section 9.6 + | "DELETE" ; Section 9.7 + | "TRACE" ; Section 9.8 + | "CONNECT" ; Section 9.9 + | extension-method + extension-method = token + token = 1* + */ + return len(method) > 0 && strings.IndexFunc(method, isNotToken) == -1 +} + +// copied from net/http/http.go +func isNotToken(r rune) bool { + return !httpguts.IsTokenRune(r) +} + +// CloseIdleConnections closes any QUIC connections in the transport's pool that are currently idle. +// An idle connection is one that was previously used for requests but is now sitting unused. +// This method does not interrupt any connections currently in use. +// It also does not affect connections obtained via NewClientConn. +func (t *Transport) CloseIdleConnections() { + t.mutex.Lock() + defer t.mutex.Unlock() + for hostname, cl := range t.clients { + if cl.useCount.Load() == 0 { + cl.Close() + delete(t.clients, hostname) + } + } +} diff --git a/third_party/quic-go/http3/transport_test.go b/third_party/quic-go/http3/transport_test.go new file mode 100644 index 0000000..4a5e095 --- /dev/null +++ b/third_party/quic-go/http3/transport_test.go @@ -0,0 +1,572 @@ +package http3 + +import ( + "bytes" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type mockBody struct { + reader bytes.Reader + readErr error + closeErr error + closed bool +} + +// make sure the mockBody can be used as a http.Request.Body +var _ io.ReadCloser = &mockBody{} + +func (m *mockBody) Read(p []byte) (int, error) { + if m.readErr != nil { + return 0, m.readErr + } + return m.reader.Read(p) +} + +func (m *mockBody) SetData(data []byte) { + m.reader = *bytes.NewReader(data) +} + +func (m *mockBody) Close() error { + m.closed = true + return m.closeErr +} + +func TestRequestValidation(t *testing.T) { + var tr Transport + + for _, tt := range []struct { + name string + req *http.Request + expectedErr string + expectedErrContains string + }{ + { + name: "plain HTTP", + req: httptest.NewRequest(http.MethodGet, "http://www.example.org/", nil), + expectedErr: "http3: unsupported protocol scheme: http", + }, + { + name: "missing URL", + req: func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://www.example.org/", nil) + r.URL = nil + return r + }(), + expectedErr: "http3: nil Request.URL", + }, + { + name: "missing URL Host", + req: func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://www.example.org/", nil) + r.URL.Host = "" + return r + }(), + expectedErr: "http3: no Host in request URL", + }, + { + name: "missing header", + req: func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://www.example.org/", nil) + r.Header = nil + return r + }(), + expectedErr: "http3: nil Request.Header", + }, + { + name: "invalid header name", + req: func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://www.example.org/", nil) + r.Header.Add("foobär", "value") + return r + }(), + expectedErr: "http3: invalid http header field name \"foobär\"", + }, + { + name: "invalid header value", + req: func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://www.example.org/", nil) + r.Header.Add("Authorization", "Bearer secret\x00") + return r + }(), + expectedErr: `http3: invalid http header field value for key "Authorization"`, + }, + { + name: "invalid method", + req: func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://www.example.org/", nil) + r.Method = "foobär" + return r + }(), + expectedErr: "http3: invalid method \"foobär\"", + }, + } { + t.Run(tt.name, func(t *testing.T) { + tt.req.Body = &mockBody{} + _, err := tr.RoundTrip(tt.req) + if tt.expectedErr != "" { + require.EqualError(t, err, tt.expectedErr) + } + if tt.expectedErrContains != "" { + require.Error(t, err) + require.Contains(t, err.Error(), tt.expectedErrContains) + } + require.True(t, tt.req.Body.(*mockBody).closed) + }) + } +} + +func TestTransportDialHostname(t *testing.T) { + type hostnameConfig struct { + dialHostname string + tlsServerName string + } + hostnameChan := make(chan hostnameConfig, 1) + tr := &Transport{ + Dial: func(_ context.Context, hostname string, tlsConf *tls.Config, _ *quic.Config) (*quic.Conn, error) { + hostnameChan <- hostnameConfig{ + dialHostname: hostname, + tlsServerName: tlsConf.ServerName, + } + return nil, errors.New("test done") + }, + } + + t.Run("port set", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net:1234", nil) + _, err := tr.RoundTripOpt(req, RoundTripOpt{}) + require.EqualError(t, err, "test done") + select { + case c := <-hostnameChan: + require.Equal(t, "quic-go.net:1234", c.dialHostname) + require.Equal(t, "quic-go.net", c.tlsServerName) + case <-time.After(1 * time.Second): + t.Fatal("timeout") + } + }) + + // if the request doesn't have a port, the default port is used + t.Run("port not set", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net", nil) + _, err := tr.RoundTripOpt(req, RoundTripOpt{}) + require.EqualError(t, err, "test done") + select { + case c := <-hostnameChan: + require.Equal(t, "quic-go.net:443", c.dialHostname) + require.Equal(t, "quic-go.net", c.tlsServerName) + case <-time.After(1 * time.Second): + t.Fatal("timeout") + } + }) +} + +func TestTransportDatagrams(t *testing.T) { + // if the default quic.Config is used, the transport automatically enables QUIC datagrams + t.Run("default quic.Config", func(t *testing.T) { + tr := &Transport{ + EnableDatagrams: true, + Dial: func(_ context.Context, _ string, _ *tls.Config, quicConf *quic.Config) (*quic.Conn, error) { + require.True(t, quicConf.EnableDatagrams) + return nil, assert.AnError + }, + } + req := httptest.NewRequest(http.MethodGet, "https://example.com", nil) + _, err := tr.RoundTripOpt(req, RoundTripOpt{}) + require.ErrorIs(t, err, assert.AnError) + }) + + // if a custom quic.Config is used, the transport just checks that QUIC datagrams are enabled + t.Run("custom quic.Config", func(t *testing.T) { + tr := &Transport{ + EnableDatagrams: true, + QUICConfig: &quic.Config{EnableDatagrams: false}, + Dial: func(_ context.Context, _ string, _ *tls.Config, quicConf *quic.Config) (*quic.Conn, error) { + t.Fatal("dial should not be called") + return nil, nil + }, + } + req := httptest.NewRequest(http.MethodGet, "https://example.com", nil) + _, err := tr.RoundTripOpt(req, RoundTripOpt{}) + require.EqualError(t, err, "HTTP Datagrams enabled, but QUIC Datagrams disabled") + }) +} + +func TestTransportMultipleQUICVersions(t *testing.T) { + qconf := &quic.Config{ + Versions: []quic.Version{quic.Version2, quic.Version1}, + } + tr := &Transport{QUICConfig: qconf} + req := httptest.NewRequest(http.MethodGet, "https://example.com", nil) + _, err := tr.RoundTrip(req) + require.EqualError(t, err, "can only use a single QUIC version for dialing a HTTP/3 connection") +} + +func TestTransportConnectionReuse(t *testing.T) { + conn, _ := newConnPair(t) + mockCtrl := gomock.NewController(t) + cl := NewMockClientConn(mockCtrl) + var dialCount int + tr := &Transport{ + Dial: func(context.Context, string, *tls.Config, *quic.Config) (*quic.Conn, error) { + dialCount++ + return conn, nil + }, + newClientConn: func(*quic.Conn) clientConn { return cl }, + } + + req1 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file1.html", nil) + // if OnlyCachedConn is set, no connection is dialed + _, err := tr.RoundTripOpt(req1, RoundTripOpt{OnlyCachedConn: true}) + require.ErrorIs(t, err, ErrNoCachedConn) + require.Zero(t, dialCount) + + // the first request establishes the connection... + cl.EXPECT().RoundTrip(req1).Return(&http.Response{Request: req1}, nil) + rsp, err := tr.RoundTrip(req1) + require.NoError(t, err) + require.Equal(t, req1, rsp.Request) + require.Equal(t, 1, dialCount) + + // ... which is then used for the second request + req2 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file2.html", nil) + cl.EXPECT().RoundTrip(req2).Return(&http.Response{Request: req2}, nil) + rsp, err = tr.RoundTrip(req2) + require.NoError(t, err) + require.Equal(t, req2, rsp.Request) + require.Equal(t, 1, dialCount) +} + +// Requests reuse the same underlying QUIC connection. +// If a request experiences an error, the behavior depends on the nature of that error. +func TestTransportConnectionRedial(t *testing.T) { + nonRetryableReq := httptest.NewRequest( + http.MethodGet, + "https://quic-go.org", + strings.NewReader("foobar"), + ) + require.Nil(t, nonRetryableReq.GetBody) + + retryableReq := nonRetryableReq.Clone(context.Background()) + retryableReq.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader("foobaz")), nil + } + + // If the error occurs when opening the stream, it is safe to retry the request: + // We can be certain that it wasn't sent out (not even partially). + t.Run("error when opening the stream", func(t *testing.T) { + require.NoError(t, + testTransportConnectionRedial(t, nonRetryableReq, &errConnUnusable{errors.New("test")}, "foobar", true), + ) + }) + + // If the error occurs when opening the stream, it is safe to retry the request: + // We can be certain that it wasn't sent out (not even partially). + t.Run("non-retryable request error after opening the stream", func(t *testing.T) { + require.ErrorIs(t, + testTransportConnectionRedial(t, nonRetryableReq, assert.AnError, "foobar", false), + assert.AnError, + ) + }) + + t.Run("retryable request after opening the stream", func(t *testing.T) { + require.ErrorIs(t, + testTransportConnectionRedial(t, retryableReq, assert.AnError, "", false), + assert.AnError, + ) + }) + + t.Run("retryable request after H3_REQUEST_REJECTED", func(t *testing.T) { + require.NoError(t, + testTransportConnectionRedial(t, + retryableReq, + &Error{ErrorCode: ErrCodeRequestRejected}, + "foobaz", + true, + ), + ) + }) + + t.Run("retryable request where GetBody returns an error", func(t *testing.T) { + req := nonRetryableReq.Clone(context.Background()) + req.GetBody = func() (io.ReadCloser, error) { + return nil, assert.AnError + } + require.ErrorIs(t, + testTransportConnectionRedial(t, req, &Error{ErrorCode: ErrCodeRequestRejected}, "", false), + assert.AnError, + ) + }) +} + +func testTransportConnectionRedial(t *testing.T, req *http.Request, roundtripErr error, expectedBody string, expectRedial bool) error { + conn, _ := newConnPair(t) + mockCtrl := gomock.NewController(t) + cl := NewMockClientConn(mockCtrl) + var dialCount int + tr := &Transport{ + Dial: func(context.Context, string, *tls.Config, *quic.Config) (*quic.Conn, error) { + dialCount++ + return conn, nil + }, + newClientConn: func(*quic.Conn) clientConn { return cl }, + } + + var body string + cl.EXPECT().RoundTrip(req).Return(nil, roundtripErr) + if expectRedial { + cl.EXPECT().RoundTrip(gomock.Any()).DoAndReturn(func(r *http.Request) (*http.Response, error) { + b, err := io.ReadAll(r.Body) + if err != nil { + panic(fmt.Sprintf("reading body failed: %v", err)) + } + body = string(b) + return &http.Response{Request: req}, nil + }) + } + + _, err := tr.RoundTrip(req) + if !expectRedial { + assert.Equal(t, 1, dialCount) + } else { + assert.Equal(t, 2, dialCount) + assert.Equal(t, expectedBody, body) + } + return err +} + +func TestTransportRequestContextCancellation(t *testing.T) { + mockCtrl := gomock.NewController(t) + cl := NewMockClientConn(mockCtrl) + conn, _ := newConnPair(t) + var dialCount int + tr := &Transport{ + Dial: func(context.Context, string, *tls.Config, *quic.Config) (*quic.Conn, error) { + dialCount++ + return conn, nil + }, + newClientConn: func(*quic.Conn) clientConn { return cl }, + } + + // the first request succeeds + req1 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file1.html", nil) + cl.EXPECT().RoundTrip(req1).Return(&http.Response{Request: req1}, nil) + rsp, err := tr.RoundTrip(req1) + require.NoError(t, err) + require.Equal(t, req1, rsp.Request) + require.Equal(t, 1, dialCount) + + // the second request reuses the QUIC connection, and runs into the cancelled context + req2 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file2.html", nil) + ctx, cancel := context.WithCancel(context.Background()) + req2 = req2.WithContext(ctx) + cl.EXPECT().RoundTrip(req2).DoAndReturn( + func(r *http.Request) (*http.Response, error) { + cancel() + return nil, context.Canceled + }, + ) + _, err = tr.RoundTrip(req2) + require.ErrorIs(t, err, context.Canceled) + require.Equal(t, 1, dialCount) + + // the next request reuses the QUIC connection + req3 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file2.html", nil) + cl.EXPECT().RoundTrip(req3).Return(&http.Response{Request: req3}, nil) + rsp, err = tr.RoundTrip(req3) + require.NoError(t, err) + require.Equal(t, req3, rsp.Request) + require.Equal(t, 1, dialCount) +} + +func TestTransportConnetionRedialHandshakeError(t *testing.T) { + mockCtrl := gomock.NewController(t) + cl := NewMockClientConn(mockCtrl) + conn, _ := newConnPair(t) + var dialCount int + tr := &Transport{ + Dial: func(context.Context, string, *tls.Config, *quic.Config) (*quic.Conn, error) { + dialCount++ + if dialCount == 1 { + return nil, assert.AnError + } + return conn, nil + }, + newClientConn: func(*quic.Conn) clientConn { return cl }, + } + + req1 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file1.html", nil) + _, err := tr.RoundTrip(req1) + require.ErrorIs(t, err, assert.AnError) + require.Equal(t, 1, dialCount) + + req2 := httptest.NewRequest(http.MethodGet, "https://quic-go.net/file2.html", nil) + cl.EXPECT().RoundTrip(req2).Return(&http.Response{Request: req2}, nil) + rsp, err := tr.RoundTrip(req2) + require.NoError(t, err) + require.Equal(t, req2, rsp.Request) + require.Equal(t, 2, dialCount) +} + +func TestTransportCloseEstablishedConnections(t *testing.T) { + mockCtrl := gomock.NewController(t) + conn, _ := newConnPair(t) + tr := &Transport{ + Dial: func(context.Context, string, *tls.Config, *quic.Config) (*quic.Conn, error) { + return conn, nil + }, + newClientConn: func(*quic.Conn) clientConn { + cl := NewMockClientConn(mockCtrl) + cl.EXPECT().RoundTrip(gomock.Any()).Return(&http.Response{}, nil) + return cl + }, + } + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net/foobar.html", nil) + _, err := tr.RoundTrip(req) + require.NoError(t, err) + require.NoError(t, tr.Close()) + + select { + case <-conn.Context().Done(): + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportCloseInFlightDials(t *testing.T) { + tr := &Transport{ + Dial: func(ctx context.Context, _ string, _ *tls.Config, _ *quic.Config) (*quic.Conn, error) { + var err error + select { + case <-ctx.Done(): + err = ctx.Err() + case <-time.After(time.Second): + err = errors.New("timeout") + } + return nil, err + }, + } + req := httptest.NewRequest(http.MethodGet, "https://quic-go.net/foobar.html", nil) + + errChan := make(chan error, 1) + go func() { + _, err := tr.RoundTrip(req) + errChan <- err + }() + + select { + case err := <-errChan: + t.Fatalf("received unexpected error: %v", err) + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + require.NoError(t, tr.Close()) + select { + case err := <-errChan: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportCloseIdleConnections(t *testing.T) { + mockCtrl := gomock.NewController(t) + conn1, _ := newConnPair(t) + conn2, _ := newConnPair(t) + roundTripCalled := make(chan struct{}) + tr := &Transport{ + Dial: func(_ context.Context, hostname string, _ *tls.Config, _ *quic.Config) (*quic.Conn, error) { + switch hostname { + case "site1.com:443": + return conn1, nil + case "site2.com:443": + return conn2, nil + default: + t.Fatal("unexpected hostname") + return nil, errors.New("unexpected hostname") + } + }, + newClientConn: func(*quic.Conn) clientConn { + cl := NewMockClientConn(mockCtrl) + cl.EXPECT().RoundTrip(gomock.Any()).DoAndReturn(func(r *http.Request) (*http.Response, error) { + roundTripCalled <- struct{}{} + <-r.Context().Done() + return nil, nil + }) + return cl + }, + } + req1 := httptest.NewRequest(http.MethodGet, "https://site1.com", nil) + req2 := httptest.NewRequest(http.MethodGet, "https://site2.com", nil) + require.NotEqual(t, req1.Host, req2.Host) + ctx1, cancel1 := context.WithCancel(context.Background()) + ctx2, cancel2 := context.WithCancel(context.Background()) + req1 = req1.WithContext(ctx1) + req2 = req2.WithContext(ctx2) + reqFinished := make(chan struct{}) + go func() { + tr.RoundTrip(req1) + reqFinished <- struct{}{} + }() + go func() { + tr.RoundTrip(req2) + reqFinished <- struct{}{} + }() + <-roundTripCalled + <-roundTripCalled + // Both two requests are started. + cancel1() + <-reqFinished + // req1 is finished + tr.CloseIdleConnections() + select { + case <-conn1.Context().Done(): + case <-time.After(time.Second): + t.Fatal("timeout") + } + + cancel2() + <-reqFinished + // all requests are finished + tr.CloseIdleConnections() + select { + case <-conn2.Context().Done(): + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportClose(t *testing.T) { + mockCtrl := gomock.NewController(t) + conn, _ := newConnPair(t) + tr := &Transport{ + Dial: func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + return conn, nil + }, + newClientConn: func(*quic.Conn) clientConn { + cl := NewMockClientConn(mockCtrl) + cl.EXPECT().RoundTrip(gomock.Any()).Return(nil, nil) + return cl + }, + } + req, err := http.NewRequest(http.MethodGet, "https://example.com", nil) + require.NoError(t, err) + _, err = tr.RoundTrip(req) + require.NoError(t, err) + require.NoError(t, tr.Close()) + _, err = tr.RoundTrip(req) + require.ErrorIs(t, err, ErrTransportClosed) +} diff --git a/third_party/quic-go/integrationtests/self/benchmark_test.go b/third_party/quic-go/integrationtests/self/benchmark_test.go new file mode 100644 index 0000000..f0a7342 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/benchmark_test.go @@ -0,0 +1,151 @@ +package self_test + +import ( + "bytes" + "context" + "fmt" + "io" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/require" +) + +func BenchmarkHandshake(b *testing.B) { + b.ReportAllocs() + + ln, err := quic.Listen(newUDPConnLocalhost(b), tlsConfig, nil) + require.NoError(b, err) + defer ln.Close() + + connChan := make(chan *quic.Conn, 1) + go func() { + for { + conn, err := ln.Accept(context.Background()) + if err != nil { + return + } + connChan <- conn + } + }() + + tr := &quic.Transport{Conn: newUDPConnLocalhost(b)} + defer tr.Close() + + for b.Loop() { + c, err := tr.Dial(context.Background(), ln.Addr(), tlsClientConfig, nil) + if err != nil { + b.Fatalf("error dialing: %v", err) + } + serverConn := <-connChan + serverConn.CloseWithError(0, "") + c.CloseWithError(0, "") + } +} + +func BenchmarkStreamChurn(b *testing.B) { + b.ReportAllocs() + + ln, err := quic.Listen(newUDPConnLocalhost(b), tlsConfig, &quic.Config{MaxIncomingStreams: 1e10}) + require.NoError(b, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(b), ln.Addr(), tlsClientConfig, nil) + require.NoError(b, err) + defer conn.CloseWithError(0, "") + + serverConn, err := ln.Accept(context.Background()) + require.NoError(b, err) + defer serverConn.CloseWithError(0, "") + + go func() { + for { + str, err := serverConn.AcceptStream(context.Background()) + if err != nil { + return + } + str.Close() + } + }() + + for b.Loop() { + str, err := conn.OpenStreamSync(context.Background()) + if err != nil { + b.Fatalf("error opening stream: %v", err) + } + if err := str.Close(); err != nil { + b.Fatalf("error closing stream: %v", err) + } + } +} + +func BenchmarkTransfer(b *testing.B) { + b.Run(fmt.Sprintf("%d kb", len(PRData)/1024), func(b *testing.B) { benchmarkTransfer(b, PRData) }) + b.Run(fmt.Sprintf("%d kb", len(PRDataLong)/1024), func(b *testing.B) { benchmarkTransfer(b, PRDataLong) }) +} + +func benchmarkTransfer(b *testing.B, data []byte) { + b.ReportAllocs() + + ln, err := quic.Listen(newUDPConnLocalhost(b), tlsConfig, nil) + require.NoError(b, err) + defer ln.Close() + + connChan := make(chan *quic.Conn, 1) + go func() { + for { + conn, err := ln.Accept(context.Background()) + if err != nil { + return + } + connChan <- conn + str, err := conn.OpenUniStream() + if err != nil { + b.Logf("error opening stream: %v", err) + return + } + if _, err := str.Write(data); err != nil { + b.Logf("error writing data: %v", err) + return + } + if err := str.Close(); err != nil { + b.Logf("error closing stream: %v", err) + return + } + } + }() + + tr := &quic.Transport{Conn: newUDPConnLocalhost(b)} + defer tr.Close() + + buf := make([]byte, len(data)) + + for b.Loop() { + c, err := tr.Dial(context.Background(), ln.Addr(), tlsClientConfig, nil) + if err != nil { + b.Fatalf("error dialing: %v", err) + } + + str, err := c.AcceptUniStream(context.Background()) + if err != nil { + b.Fatalf("error accepting stream: %v", err) + } + if _, err := io.ReadFull(str, buf); err != nil { + b.Fatalf("error reading data: %v", err) + } + if _, err := str.Read([]byte{0}); err != io.EOF { + b.Fatalf("error reading EOF: %v", err) + } + if !bytes.Equal(buf, data) { + b.Fatalf("data mismatch: got %x, expected %x", buf, data) + } + + serverConn := <-connChan + serverConn.CloseWithError(0, "") + c.CloseWithError(0, "") + } +} diff --git a/third_party/quic-go/integrationtests/self/cancelation_test.go b/third_party/quic-go/integrationtests/self/cancelation_test.go new file mode 100644 index 0000000..f9d00c2 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/cancelation_test.go @@ -0,0 +1,598 @@ +package self_test + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "math/rand/v2" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestStreamReadCancellation(t *testing.T) { + t.Run("immediate", func(t *testing.T) { + testStreamCancellation(t, func(str *quic.ReceiveStream) error { + str.CancelRead(quic.StreamErrorCode(str.StreamID())) + _, err := str.Read([]byte{0}) + return err + }, nil) + }) + + t.Run("after reading some data", func(t *testing.T) { + testStreamCancellation(t, func(str *quic.ReceiveStream) error { + length := rand.IntN(len(PRData) - 1) + if _, err := io.ReadAll(io.LimitReader(str, int64(length))); err != nil { + return fmt.Errorf("reading stream data failed: %w", err) + } + str.CancelRead(quic.StreamErrorCode(str.StreamID())) + _, err := str.Read([]byte{0}) + return err + }, nil) + }) + + // This test is especially valuable when run with race detector, + // see https://github.com/apernet/quic-go/issues/3239. + t.Run("concurrent", func(t *testing.T) { + testStreamCancellation(t, func(str *quic.ReceiveStream) error { + errChan := make(chan error, 1) + go func() { + for { + if _, err := str.Read(make([]byte, 16)); err != nil { + errChan <- err + return + } + time.Sleep(time.Millisecond) + } + }() + + done := make(chan struct{}) + go func() { + defer close(done) + str.CancelRead(quic.StreamErrorCode(str.StreamID())) + }() + + timeout := time.After(time.Second) + select { + case <-done: + case <-timeout: + return fmt.Errorf("timeout canceling") + } + select { + case err := <-errChan: + return err + case <-timeout: + return fmt.Errorf("timeout canceling") + } + }, nil) + }) +} + +func TestStreamWriteCancellation(t *testing.T) { + t.Run("immediate", func(t *testing.T) { + testStreamCancellation(t, nil, func(str *quic.SendStream) error { + str.CancelWrite(quic.StreamErrorCode(str.StreamID())) + _, err := str.Write([]byte{0}) + return err + }) + }) + + t.Run("after writing some data", func(t *testing.T) { + testStreamCancellation(t, nil, func(str *quic.SendStream) error { + length := rand.IntN(len(PRData) - 1) + if _, err := str.Write(PRData[:length]); err != nil { + return fmt.Errorf("writing stream data failed: %w", err) + } + str.CancelWrite(quic.StreamErrorCode(str.StreamID())) + _, err := str.Write([]byte{0}) + return err + }) + }) + + // This test is especially valuable when run with race detector, + // see https://github.com/apernet/quic-go/issues/3239. + t.Run("concurrent", func(t *testing.T) { + testStreamCancellation(t, nil, func(str *quic.SendStream) error { + errChan := make(chan error, 1) + go func() { + var offset int + for { + n, err := str.Write(PRData[offset : offset+128]) + if err != nil { + errChan <- err + return + } + offset += n + time.Sleep(time.Millisecond) + } + }() + + done := make(chan struct{}) + go func() { + defer close(done) + str.CancelWrite(quic.StreamErrorCode(str.StreamID())) + }() + + timeout := time.After(time.Second) + select { + case <-done: + case <-timeout: + return fmt.Errorf("timeout canceling") + } + select { + case err := <-errChan: + return err + case <-timeout: + return fmt.Errorf("timeout canceling") + } + }) + }) +} + +func TestStreamReadWriteCancellation(t *testing.T) { + t.Run("immediate", func(t *testing.T) { + testStreamCancellation(t, + func(str *quic.ReceiveStream) error { + str.CancelRead(quic.StreamErrorCode(str.StreamID())) + _, err := str.Read([]byte{0}) + return err + }, + func(str *quic.SendStream) error { + str.CancelWrite(quic.StreamErrorCode(str.StreamID())) + _, err := str.Write([]byte{0}) + return err + }, + ) + }) + + t.Run("after writing some data", func(t *testing.T) { + testStreamCancellation(t, + func(str *quic.ReceiveStream) error { + length := rand.IntN(len(PRData) - 1) + if _, err := io.ReadAll(io.LimitReader(str, int64(length))); err != nil { + return fmt.Errorf("reading stream data failed: %w", err) + } + str.CancelRead(quic.StreamErrorCode(str.StreamID())) + _, err := str.Read([]byte{0}) + return err + }, + func(str *quic.SendStream) error { + length := rand.IntN(len(PRData) - 1) + if _, err := str.Write(PRData[:length]); err != nil { + return fmt.Errorf("writing stream data failed: %w", err) + } + str.CancelWrite(quic.StreamErrorCode(str.StreamID())) + _, err := str.Write([]byte{0}) + return err + }, + ) + }) +} + +// If readFunc is set, the read side is canceled for 50% of the streams. +// If writeFunc is set, the write side is canceled for 50% of the streams. +func testStreamCancellation( + t *testing.T, + readFunc func(str *quic.ReceiveStream) error, + writeFunc func(str *quic.SendStream) error, +) { + const numStreams = 80 + + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(2*time.Second)) + defer cancel() + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIncomingUniStreams: numStreams / 2}), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + + type cancellationErr struct { + StreamID quic.StreamID + Err error + } + + var numCancellations int + actions := make([]bool, numStreams) + for i := range actions { + actions[i] = rand.IntN(2) == 0 + if actions[i] { + numCancellations++ + } + } + + // The server accepts a single connection, and then opens numStreams unidirectional streams. + // On each of these streams, it (tries to) write PRData. + serverErrChan := make(chan *cancellationErr, numStreams) + go func() { + for _, doCancel := range actions { + str, err := serverConn.OpenUniStreamSync(ctx) + if err != nil { + serverErrChan <- &cancellationErr{StreamID: protocol.InvalidStreamID, Err: fmt.Errorf("opening stream failed: %w", err)} + return + } + go func() { + if writeFunc != nil && doCancel { + if err := writeFunc(str); err != nil { + serverErrChan <- &cancellationErr{StreamID: str.StreamID(), Err: err} + return + } + serverErrChan <- nil + return + } + defer str.Close() + if _, err := str.Write(PRData); err != nil { + serverErrChan <- &cancellationErr{StreamID: str.StreamID(), Err: err} + return + } + serverErrChan <- nil + }() + } + }() + + clientErrChan := make(chan *cancellationErr, numStreams) + for _, doCancel := range actions { + str, err := conn.AcceptUniStream(ctx) + require.NoError(t, err) + go func(str *quic.ReceiveStream) { + if readFunc != nil && doCancel { + if err := readFunc(str); err != nil { + clientErrChan <- &cancellationErr{StreamID: str.StreamID(), Err: err} + return + } + } + data, err := io.ReadAll(str) + if err != nil { + clientErrChan <- &cancellationErr{StreamID: str.StreamID(), Err: fmt.Errorf("reading stream data failed: %w", err)} + return + } + if !bytes.Equal(data, PRData) { + clientErrChan <- &cancellationErr{StreamID: str.StreamID(), Err: fmt.Errorf("received data mismatch")} + return + } + clientErrChan <- nil + }(str) + } + + timeout := time.After(time.Second) + var clientErrs, serverErrs int + for range numStreams { + select { + case err := <-serverErrChan: + if err != nil { + if err.StreamID == protocol.InvalidStreamID { // failed opening a stream + require.NoError(t, err.Err) + continue + } + var streamErr *quic.StreamError + require.ErrorAs(t, err.Err, &streamErr) + assert.Equal(t, streamErr.StreamID, err.StreamID) + assert.Equal(t, streamErr.ErrorCode, quic.StreamErrorCode(err.StreamID)) + if readFunc != nil && writeFunc == nil { + assert.Equal(t, streamErr.Remote, readFunc != nil) + } + serverErrs++ + } + case <-timeout: + t.Fatalf("timeout") + } + select { + case err := <-clientErrChan: + if err != nil { + if err.StreamID == protocol.InvalidStreamID { // failed accepting a stream + require.NoError(t, err.Err) + continue + } + var streamErr *quic.StreamError + require.ErrorAs(t, err.Err, &streamErr) + assert.Equal(t, streamErr.StreamID, err.StreamID) + assert.Equal(t, streamErr.ErrorCode, quic.StreamErrorCode(err.StreamID)) + if readFunc != nil && writeFunc == nil { + assert.Equal(t, streamErr.Remote, writeFunc != nil) + } + clientErrs++ + } + case <-timeout: + t.Fatalf("timeout") + } + } + assert.Equal(t, numCancellations, clientErrs, "client canceled streams") + // The server will only count a stream as being reset if it learns about the cancellation + // before it finished writing all data. + assert.LessOrEqual(t, serverErrs, numCancellations, "server-observed canceled streams") + assert.NotZero(t, serverErrs, "server-observed canceled streams") +} + +func TestCancelAcceptStream(t *testing.T) { + const numStreams = 30 + + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIncomingUniStreams: numStreams / 3}), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverErrChan := make(chan error, 1) + go func() { + defer close(serverErrChan) + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(2*time.Second)) + defer cancel() + ticker := time.NewTicker(5 * time.Millisecond) + defer ticker.Stop() + for range numStreams { + <-ticker.C + str, err := serverConn.OpenUniStreamSync(ctx) + if err != nil { + serverErrChan <- err + return + } + if _, err := str.Write(PRData); err != nil { + serverErrChan <- err + return + } + str.Close() + } + }() + + var numToAccept int + var counter atomic.Int32 + var wg sync.WaitGroup + wg.Add(numStreams) + for numToAccept < numStreams { + ctx, cancel := context.WithCancel(context.Background()) + // cancel accepting half of the streams + if rand.Int()%2 == 0 { + cancel() + } else { + numToAccept++ + defer cancel() + } + + go func() { + str, err := conn.AcceptUniStream(ctx) + if err != nil { + if errors.Is(err, context.Canceled) { + counter.Add(1) + } + return + } + go func() { + data, err := io.ReadAll(str) + if err != nil { + t.Errorf("ReadAll failed: %v", err) + return + } + if !bytes.Equal(data, PRData) { + t.Errorf("received data mismatch") + return + } + wg.Done() + }() + }() + } + wg.Wait() + + count := counter.Load() + t.Logf("canceled AcceptStream %d times", count) + require.Greater(t, count, int32(numStreams/4)) + require.NoError(t, conn.CloseWithError(0, "")) + require.NoError(t, server.Close()) + require.NoError(t, <-serverErrChan) +} + +func TestCancelOpenStreamSync(t *testing.T) { + const ( + numStreams = 16 + maxIncomingStreams = 4 + ) + + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + conn, err := quic.Dial( + context.Background(), + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIncomingUniStreams: maxIncomingStreams}), + ) + require.NoError(t, err) + + msg := make(chan struct{}, 1) + serverErrChan := make(chan error, numStreams+1) + var numCanceled int + serverConn, err := server.Accept(context.Background()) + require.NoError(t, err) + go func() { + defer close(msg) + var numOpened int + for numOpened < numStreams { + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(10*time.Millisecond)) + defer cancel() + str, err := serverConn.OpenUniStreamSync(ctx) + if err != nil { + if !errors.Is(err, context.DeadlineExceeded) { + serverErrChan <- err + return + } + numCanceled++ + select { + case msg <- struct{}{}: + default: + } + continue + } + numOpened++ + go func(str *quic.SendStream) { + defer str.Close() + if _, err := str.Write(PRData); err != nil { + serverErrChan <- err + } + }(str) + } + }() + + clientErrChan := make(chan error, numStreams) + for range numStreams { + <-msg + str, err := conn.AcceptUniStream(context.Background()) + require.NoError(t, err) + go func(str *quic.ReceiveStream) { + data, err := io.ReadAll(str) + if err != nil { + clientErrChan <- err + return + } + if !bytes.Equal(data, PRData) { + clientErrChan <- fmt.Errorf("received data mismatch") + return + } + clientErrChan <- nil + }(str) + } + + timeout := time.After(scaleDuration(2 * time.Second)) + for range numStreams { + select { + case err := <-clientErrChan: + require.NoError(t, err) + case err := <-serverErrChan: + require.NoError(t, err) + case <-timeout: + t.Fatalf("timeout") + } + } + + count := numCanceled + t.Logf("Canceled OpenStreamSync %d times", count) + require.GreaterOrEqual(t, count, numStreams-maxIncomingStreams) + require.NoError(t, conn.CloseWithError(0, "")) + require.NoError(t, server.Close()) +} + +func TestHeavyStreamCancellation(t *testing.T) { + const maxIncomingStreams = 500 + + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{MaxIncomingStreams: maxIncomingStreams, MaxIdleTimeout: 10 * time.Second}), + ) + require.NoError(t, err) + defer server.Close() + + var wg sync.WaitGroup + wg.Add(2 * 4 * maxIncomingStreams) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + serverConn, err := server.Accept(context.Background()) + require.NoError(t, err) + + handleStream := func(str *quic.Stream) { + str.SetDeadline(time.Now().Add(time.Second)) + go func() { + defer wg.Done() + if rand.Int()%2 == 0 { + io.ReadAll(str) + } + }() + go func() { + defer wg.Done() + if rand.Int()%2 == 0 { + str.Write([]byte("foobar")) + if rand.Int()%2 == 0 { + str.Close() + } + } + }() + go func() { + defer wg.Done() + // Make sure we at least send out *something* for the last stream, + // otherwise the peer might never receive this anything for this stream. + if rand.Int()%2 == 0 || str.StreamID() == 4*(maxIncomingStreams-1) { + str.CancelWrite(1234) + } + }() + go func() { + defer wg.Done() + if rand.Int()%2 == 0 { + str.CancelRead(1234) + } + }() + } + + serverErrChan := make(chan error, 1) + go func() { + defer close(serverErrChan) + + for { + str, err := serverConn.AcceptStream(context.Background()) + if err != nil { + serverErrChan <- err + return + } + handleStream(str) + } + }() + + for range maxIncomingStreams { + str, err := conn.OpenStreamSync(context.Background()) + require.NoError(t, err) + handleStream(str) + } + + // We don't expect to accept any stream here. + // We're just making sure the connection stays open and there's no error. + ctx, cancel = context.WithTimeout(context.Background(), scaleDuration(50*time.Millisecond)) + defer cancel() + _, err = conn.AcceptStream(ctx) + require.ErrorIs(t, err, context.DeadlineExceeded) + + wg.Wait() + + require.NoError(t, conn.CloseWithError(0, "")) + select { + case err := <-serverErrChan: + require.IsType(t, &quic.ApplicationError{}, err) + case <-time.After(scaleDuration(time.Second)): + t.Fatal("timeout waiting for server to stop") + } +} diff --git a/third_party/quic-go/integrationtests/self/chrome_parrot_test.go b/third_party/quic-go/integrationtests/self/chrome_parrot_test.go new file mode 100644 index 0000000..f08bdff --- /dev/null +++ b/third_party/quic-go/integrationtests/self/chrome_parrot_test.go @@ -0,0 +1,195 @@ +package self_test + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "io" + "math/big" + "net" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/stretchr/testify/require" +) + +// chromeCompatibleTLSConfigs returns a server/client config pair using an ECDSA +// P-256 certificate. +// +// The shared test helpers use Ed25519, which a Chrome-parroting client cannot +// verify: its signature_algorithms extension advertises no Ed25519 scheme, so +// such a server has nothing to sign with and aborts. That is faithful behavior +// rather than a defect, but it is a deployment constraint; +// TestChromeParrotRejectsEd25519Server pins it down. +func chromeCompatibleTLSConfigs(t *testing.T) (server, client *tls.Config) { + t.Helper() + + caKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + caTmpl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "chrome-parrot test CA"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + IsCA: true, + KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature, + BasicConstraintsValid: true, + } + caDER, err := x509.CreateCertificate(rand.Reader, caTmpl, caTmpl, &caKey.PublicKey, caKey) + require.NoError(t, err) + ca, err := x509.ParseCertificate(caDER) + require.NoError(t, err) + + leafKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + leafTmpl := &x509.Certificate{ + SerialNumber: big.NewInt(2), + DNSNames: []string{"localhost"}, + IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.IPv6loopback}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + } + leafDER, err := x509.CreateCertificate(rand.Reader, leafTmpl, ca, &leafKey.PublicKey, caKey) + require.NoError(t, err) + + pool := x509.NewCertPool() + pool.AddCert(ca) + + const alpn = "h3" // required for an exact match + return &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: [][]byte{leafDER}, + PrivateKey: leafKey, + }}, + NextProtos: []string{alpn}, + }, &tls.Config{ + RootCAs: pool, + ServerName: "localhost", + NextProtos: []string{alpn}, + } +} + +// TestChromeParrotHandshake is the load-bearing check on Config.ChromeParrot: the +// uTLS ClientHello, the chaos-protected Initial packets and the reshaped +// transport parameters must still produce a working connection. Cosmetics that +// break interop are worse than no cosmetics. +func TestChromeParrotHandshake(t *testing.T) { + serverConf, clientConf := chromeCompatibleTLSConfigs(t) + + server, err := quic.Listen(newUDPConnLocalhost(t), serverConf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + clientConn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), + clientConf, getQuicConfig(&quic.Config{ChromeParrot: true})) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + // Only TLS 1.3 suites and an ML-KEM key share are offered, so confirm we + // actually negotiated TLS 1.3 over that. + state := clientConn.ConnectionState().TLS + require.Equal(t, uint16(tls.VersionTLS13), state.Version) + require.True(t, state.HandshakeComplete) + + // Push enough data to cross the shredded-CRYPTO reassembly path and exercise + // the pinned flow control windows. + const payloadLen = 256 << 10 + go func() { + str, err := serverConn.OpenUniStreamSync(ctx) + if err != nil { + return + } + defer str.Close() + str.Write(PRDataLong[:payloadLen]) + }() + + str, err := clientConn.AcceptUniStream(ctx) + require.NoError(t, err) + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRDataLong[:payloadLen], data) +} + +// TestChromeParrotRejectsEd25519Server documents a deployment constraint: with no +// Ed25519 signature scheme advertised, a ChromeParrot client cannot handshake +// against a server presenting an Ed25519 certificate. If this test starts +// passing, the ClientHello has drifted. +func TestChromeParrotRejectsEd25519Server(t *testing.T) { + // The shared helpers issue Ed25519 certificates. + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + _, err = quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), + getTLSClientConfig(), getQuicConfig(&quic.Config{ChromeParrot: true})) + require.Error(t, err) + require.Contains(t, err.Error(), "handshake failure") +} + +// TestChromeParrotAppliesChromeTransportParameters checks that the peer receives +// the pinned values rather than quic-go's, since what goes on the wire is the +// whole point. +func TestChromeParrotAppliesChromeTransportParameters(t *testing.T) { + serverConf, clientConf := chromeCompatibleTLSConfigs(t) + + server, err := quic.Listen(newUDPConnLocalhost(t), serverConf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + clientConn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), + clientConf, getQuicConfig(&quic.Config{ChromeParrot: true})) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + // The server may open up to the advertised initial_max_streams_bidi. quic-go's + // own default differs, so this passes only if the pinned values reached the + // wire and were honored. + for i := range chromeParrotBidiStreamLimit { + str, err := serverConn.OpenStream() + require.NoError(t, err, "opening stream %d within the advertised limit", i) + require.NotNil(t, str) + } +} + +// The advertised initial_max_streams_bidi. +const chromeParrotBidiStreamLimit = 100 + +// TestChromeParrotRejectsUnsupportedTLSConfig checks that config fields which +// cannot be carried across to uTLS fail loudly instead of being silently dropped. +// For a verification callback, silently dropping it would be a security bug. +func TestChromeParrotRejectsUnsupportedTLSConfig(t *testing.T) { + _, clientConf := chromeCompatibleTLSConfigs(t) + clientConf.VerifyConnection = func(tls.ConnectionState) error { return nil } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + _, err := quic.Dial(ctx, newUDPConnLocalhost(t), &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 1234}, + clientConf, getQuicConfig(&quic.Config{ChromeParrot: true})) + require.Error(t, err) + require.Contains(t, err.Error(), "VerifyConnection is not supported") +} diff --git a/third_party/quic-go/integrationtests/self/close_test.go b/third_party/quic-go/integrationtests/self/close_test.go new file mode 100644 index 0000000..e2be9e0 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/close_test.go @@ -0,0 +1,230 @@ +package self_test + +import ( + "context" + "crypto/tls" + "net" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestConnectionCloseRetransmission(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + serverAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002} + + var drop atomic.Bool + var mx sync.Mutex + var dropped [][]byte + n := &simnet.Simnet{ + Router: &droppingRouter{Drop: func(p simnet.Packet) bool { + shouldDrop := drop.Load() && p.From.String() == serverAddr.String() + if shouldDrop { + mx.Lock() + dropped = append(dropped, p.Data) + mx.Unlock() + } + return shouldDrop + }}, + } + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + clientConn := n.NewEndpoint(&net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001}, settings) + serverConn := n.NewEndpoint(serverAddr, settings) + require.NoError(t, n.Start()) + defer n.Close() + + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + server, err := tr.Listen( + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, clientConn, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + sconn, err := server.Accept(ctx) + require.NoError(t, err) + + time.Sleep(rtt) + + drop.Store(true) + sconn.CloseWithError(1337, "closing") + + // send 100 packets + for range 100 { + str, err := conn.OpenStream() + require.NoError(t, err) + _, err = str.Write([]byte("foobar")) + require.NoError(t, err) + + // A closed connection will drop packets if a very short queue overflows. + // Waiting for one nanosecond makes synctest process the packet before advancing + // the synthetic clock. + time.Sleep(time.Nanosecond) + } + + time.Sleep(rtt) + + mx.Lock() + defer mx.Unlock() + + // Expect retransmissions of the CONNECTION_CLOSE for the + // 1st, 2nd, 4th, 8th, 16th, 32th, 64th packet: 7 in total (+1 for the original packet) + require.Len(t, dropped, 8) + + // verify all retransmitted packets were identical + for i := 1; i < len(dropped); i++ { + require.Equal(t, dropped[0], dropped[i]) + } + }) +} + +func TestDrainServerAcceptQueue(t *testing.T) { + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + dialer := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer dialer.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + // fill up the accept queue + conns := make([]*quic.Conn, 0, protocol.MaxAcceptQueueSize) + for range protocol.MaxAcceptQueueSize { + conn, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + conns = append(conns, conn) + } + time.Sleep(scaleDuration(25 * time.Millisecond)) // wait for connections to be queued + + server.Close() + for i := range protocol.MaxAcceptQueueSize { + c, err := server.Accept(ctx) + require.NoError(t, err) + // make sure the connection is not closed + require.NoError(t, context.Cause(conns[i].Context()), "client connection closed") + require.NoError(t, context.Cause(c.Context()), "server connection closed") + c.CloseWithError(0, "") + } + _, err = server.Accept(ctx) + require.ErrorIs(t, err, quic.ErrServerClosed) +} + +type brokenConn struct { + net.PacketConn + + broken chan struct{} + breakErr atomic.Pointer[error] +} + +func newBrokenConn(conn net.PacketConn) *brokenConn { + c := &brokenConn{ + PacketConn: conn, + broken: make(chan struct{}), + } + go func() { + <-c.broken + // make calls to ReadFrom return + c.SetDeadline(time.Now()) + }() + return c +} + +func (c *brokenConn) ReadFrom(b []byte) (int, net.Addr, error) { + if err := c.breakErr.Load(); err != nil { + return 0, nil, *err + } + n, addr, err := c.PacketConn.ReadFrom(b) + if err != nil { + select { + case <-c.broken: + err = *c.breakErr.Load() + default: + } + } + return n, addr, err +} + +func (c *brokenConn) Break(e error) { + c.breakErr.Store(&e) + close(c.broken) +} + +func TestTransportClose(t *testing.T) { + t.Run("Close", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + testTransportClose(t, conn, func() { conn.Close() }, nil) + }) + + t.Run("connection error", func(t *testing.T) { + t.Setenv("QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING", "true") + + bc := newBrokenConn(newUDPConnLocalhost(t)) + testTransportClose(t, bc, func() { bc.Break(assert.AnError) }, assert.AnError) + }) +} + +func testTransportClose(t *testing.T, conn net.PacketConn, closeFn func(), expectedErr error) { + server := newUDPConnLocalhost(t) + tr := &quic.Transport{Conn: conn} + + errChan := make(chan error, 1) + go func() { + _, err := tr.Dial(context.Background(), server.LocalAddr(), &tls.Config{InsecureSkipVerify: true}, getQuicConfig(nil)) + errChan <- err + }() + + select { + case <-errChan: + t.Fatal("didn't expect Dial to return yet") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + closeFn() + + select { + case err := <-errChan: + require.Error(t, err) + require.ErrorIs(t, err, quic.ErrTransportClosed) + if expectedErr != nil { + require.ErrorIs(t, err, expectedErr) + } + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // it's not possible to dial new connections + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(50*time.Millisecond)) + defer cancel() + _, err := tr.Dial(ctx, server.LocalAddr(), &tls.Config{InsecureSkipVerify: true}, getQuicConfig(nil)) + require.Error(t, err) + require.ErrorIs(t, err, quic.ErrTransportClosed) + if expectedErr != nil { + require.ErrorIs(t, err, expectedErr) + } + + // it's not possible to create new listeners + _, err = tr.Listen(&tls.Config{}, nil) + require.Error(t, err) + require.ErrorIs(t, err, quic.ErrTransportClosed) + if expectedErr != nil { + require.ErrorIs(t, err, expectedErr) + } +} diff --git a/third_party/quic-go/integrationtests/self/conn_id_test.go b/third_party/quic-go/integrationtests/self/conn_id_test.go new file mode 100644 index 0000000..9de31ba --- /dev/null +++ b/third_party/quic-go/integrationtests/self/conn_id_test.go @@ -0,0 +1,151 @@ +package self_test + +import ( + "context" + "crypto/rand" + "fmt" + "io" + mrand "math/rand/v2" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type connIDGenerator struct { + Length int +} + +var _ quic.ConnectionIDGenerator = &connIDGenerator{} + +func (c *connIDGenerator) GenerateConnectionID() (quic.ConnectionID, error) { + b := make([]byte, c.Length) + if _, err := rand.Read(b); err != nil { + return quic.ConnectionID{}, fmt.Errorf("generating conn ID failed: %w", err) + } + return protocol.ParseConnectionID(b), nil +} + +func (c *connIDGenerator) ConnectionIDLen() int { return c.Length } + +func randomConnIDLen() int { return 2 + mrand.IntN(19) } + +func TestConnectionIDsZeroLength(t *testing.T) { + testTransferWithConnectionIDs(t, randomConnIDLen(), 0, nil, nil) +} + +func TestConnectionIDsRandomLengths(t *testing.T) { + testTransferWithConnectionIDs(t, randomConnIDLen(), randomConnIDLen(), nil, nil) +} + +func TestConnectionIDsCustomGenerator(t *testing.T) { + testTransferWithConnectionIDs(t, 0, 0, + &connIDGenerator{Length: randomConnIDLen()}, + &connIDGenerator{Length: randomConnIDLen()}, + ) +} + +// connIDLen is ignored when connIDGenerator is set +func testTransferWithConnectionIDs( + t *testing.T, + serverConnIDLen, clientConnIDLen int, + serverConnIDGenerator, clientConnIDGenerator quic.ConnectionIDGenerator, +) { + t.Helper() + + if serverConnIDGenerator != nil { + t.Logf("using %d byte connection ID generator for the server", serverConnIDGenerator.ConnectionIDLen()) + } else { + t.Logf("issuing %d byte connection ID from the server", serverConnIDLen) + } + if clientConnIDGenerator != nil { + t.Logf("using %d byte connection ID generator for the client", clientConnIDGenerator.ConnectionIDLen()) + } else { + t.Logf("issuing %d byte connection ID from the client", clientConnIDLen) + } + + // setup server + serverTr := &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnectionIDLength: serverConnIDLen, + ConnectionIDGenerator: serverConnIDGenerator, + } + defer serverTr.Close() + addTracer(serverTr) + serverCounter, serverTracer := newPacketTracer() + ln, err := serverTr.Listen( + getTLSConfig(), + getQuicConfig(&quic.Config{ + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return serverTracer + }, + }), + ) + require.NoError(t, err) + + // setup client + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + var conn *quic.Conn + clientCounter, clientTracer := newPacketTracer() + clientQUICConf := getQuicConfig(&quic.Config{ + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return clientTracer }, + }) + if clientConnIDGenerator == nil && clientConnIDLen == 0 { + conn, err = quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), clientQUICConf) + require.NoError(t, err) + } else { + clientTr := &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnectionIDLength: clientConnIDLen, + ConnectionIDGenerator: clientConnIDGenerator, + } + defer clientTr.Close() + addTracer(clientTr) + conn, err = clientTr.Dial(ctx, ln.Addr(), getTLSClientConfig(), clientQUICConf) + require.NoError(t, err) + } + + serverConn, err := ln.Accept(context.Background()) + require.NoError(t, err) + serverStr, err := serverConn.OpenStream() + require.NoError(t, err) + + go func() { + serverStr.Write(PRData) + serverStr.Close() + }() + + str, err := conn.AcceptStream(context.Background()) + require.NoError(t, err) + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRData, data) + + conn.CloseWithError(0, "") + serverConn.CloseWithError(0, "") + + for _, p := range serverCounter.getRcvdShortHeaderPackets() { + expectedLen := serverConnIDLen + if serverConnIDGenerator != nil { + expectedLen = serverConnIDGenerator.ConnectionIDLen() + } + if !assert.Equal(t, expectedLen, p.hdr.DestConnectionID.Len(), "server conn length mismatch") { + break + } + } + for _, p := range clientCounter.getRcvdShortHeaderPackets() { + expectedLen := clientConnIDLen + if clientConnIDGenerator != nil { + expectedLen = clientConnIDGenerator.ConnectionIDLen() + } + if !assert.Equal(t, expectedLen, p.hdr.DestConnectionID.Len(), "client conn length mismatch") { + break + } + } +} diff --git a/third_party/quic-go/integrationtests/self/connection_migration_test.go b/third_party/quic-go/integrationtests/self/connection_migration_test.go new file mode 100644 index 0000000..f6942aa --- /dev/null +++ b/third_party/quic-go/integrationtests/self/connection_migration_test.go @@ -0,0 +1,144 @@ +package self_test + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + + "github.com/stretchr/testify/require" +) + +func TestConnectionMigration(t *testing.T) { + ln, err := quic.ListenAddr("localhost:0", getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + tr1 := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer tr1.Close() + tr2 := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer tr2.Close() + + var packetsPath1, packetsPath2 atomic.Int64 + + const rtt = 5 * time.Millisecond + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(dir quicproxy.Direction, from, to net.Addr, _ []byte) time.Duration { + var port int + switch dir { + case quicproxy.DirectionIncoming: + port = from.(*net.UDPAddr).Port + case quicproxy.DirectionOutgoing: + port = to.(*net.UDPAddr).Port + } + switch port { + case tr1.Conn.LocalAddr().(*net.UDPAddr).Port: + packetsPath1.Add(1) + case tr2.Conn.LocalAddr().(*net.UDPAddr).Port: + packetsPath2.Add(1) + default: + fmt.Println("address not found", from) + } + return rtt / 2 + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + conn, err := tr1.Dial(ctx, proxy.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + sconn, err := ln.Accept(ctx) + require.NoError(t, err) + defer sconn.CloseWithError(0, "") + + sendAndReceiveFile := func(t *testing.T) { + t.Helper() + str, err := conn.OpenUniStream() + require.NoError(t, err) + + errChan := make(chan error, 1) + go func() { + defer close(errChan) + sstr, err := sconn.AcceptUniStream(ctx) + if err != nil { + errChan <- fmt.Errorf("accepting stream: %w", err) + return + } + data, err := io.ReadAll(sstr) + if err != nil { + errChan <- fmt.Errorf("reading stream data: %w", err) + return + } + if !bytes.Equal(data, PRData) { + errChan <- errors.New("unexpected data") + } + }() + + _, err = str.Write(PRData) + require.NoError(t, err) + require.NoError(t, str.Close()) + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timed out waiting for data") + } + } + + sendAndReceiveFile(t) // stream 2 + require.NotZero(t, packetsPath1.Load()) + require.Zero(t, packetsPath2.Load()) + + // probing the path causes a few packets to be sent on path 2 + path, err := conn.AddPath(tr2) + require.NoError(t, err) + require.ErrorIs(t, path.Switch(), quic.ErrPathNotValidated) + require.NoError(t, path.Probe(ctx)) + require.Less(t, int(packetsPath2.Load()), 5) + + // make sure that no more packets are sent on path 2 before switching to the path + c2 := packetsPath2.Load() + sendAndReceiveFile(t) // stream 6 + require.Equal(t, packetsPath2.Load(), c2) + + time.Sleep(3 * rtt) // wait for ACKs + + // now switch and make sure that no packets are sent on path 1 + require.NoError(t, path.Switch()) + sendAndReceiveFile(t) // stream 10 + c1 := packetsPath1.Load() + require.Equal(t, c1, packetsPath1.Load()) + require.Greater(t, packetsPath2.Load(), c2) + require.Equal(t, tr2.Conn.LocalAddr(), conn.LocalAddr()) + + // switch back to the handshake path + time.Sleep(3 * rtt) // wait for ACKs + c1BeforeSwitch := packetsPath1.Load() + c2BeforeSwitch := packetsPath2.Load() + path2, err := conn.AddPath(tr1) + require.NoError(t, err) + require.NoError(t, path2.Probe(ctx)) + time.Sleep(3 * rtt) // wait for ACKs + require.NoError(t, path2.Switch()) + sendAndReceiveFile(t) // stream 14 + require.Greater(t, packetsPath1.Load(), c1BeforeSwitch) + // some path probing might have happened + require.Less(t, int(packetsPath2.Load()-c2BeforeSwitch), 20) + require.Equal(t, tr1.Conn.LocalAddr(), conn.LocalAddr()) +} diff --git a/third_party/quic-go/integrationtests/self/datagram_test.go b/third_party/quic-go/integrationtests/self/datagram_test.go new file mode 100644 index 0000000..2981e72 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/datagram_test.go @@ -0,0 +1,432 @@ +package self_test + +import ( + "bytes" + "context" + "io" + mrand "math/rand/v2" + "net" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/testutils/events" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDatagramNegotiation(t *testing.T) { + t.Run("server enable, client enable", func(t *testing.T) { + testDatagramNegotiation(t, true, true) + }) + t.Run("server enable, client disable", func(t *testing.T) { + testDatagramNegotiation(t, true, false) + }) + t.Run("server disable, client enable", func(t *testing.T) { + testDatagramNegotiation(t, false, true) + }) + t.Run("server disable, client disable", func(t *testing.T) { + testDatagramNegotiation(t, false, false) + }) +} + +func TestDatagramNegotiationWithOmittedClientTransportParameter(t *testing.T) { + const maxDatagramFrameSize = 1200 + + t.Run("server assumes omitted client support", func(t *testing.T) { + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + EnableDatagrams: true, + MaxDatagramFrameSize: maxDatagramFrameSize, + AssumePeerMaxDatagramFrameSize: maxDatagramFrameSize, + }), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientConn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + EnableDatagrams: true, + MaxDatagramFrameSize: maxDatagramFrameSize, + OmitMaxDatagramFrameSize: true, + }), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + require.Equal(t, struct{ Remote, Local bool }{Remote: true, Local: true}, serverConn.ConnectionState().SupportsDatagrams) + require.Equal(t, struct{ Remote, Local bool }{Remote: true, Local: true}, clientConn.ConnectionState().SupportsDatagrams) + + require.NoError(t, serverConn.SendDatagram([]byte("foo"))) + datagram, err := clientConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("foo"), datagram) + + require.NoError(t, clientConn.SendDatagram([]byte("bar"))) + datagram, err = serverConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("bar"), datagram) + }) + + t.Run("standard server rejects omitted client support", func(t *testing.T) { + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + EnableDatagrams: true, + MaxDatagramFrameSize: maxDatagramFrameSize, + }), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientConn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + EnableDatagrams: true, + MaxDatagramFrameSize: maxDatagramFrameSize, + OmitMaxDatagramFrameSize: true, + }), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + require.Equal(t, struct{ Remote, Local bool }{Remote: false, Local: true}, serverConn.ConnectionState().SupportsDatagrams) + require.Equal(t, struct{ Remote, Local bool }{Remote: true, Local: true}, clientConn.ConnectionState().SupportsDatagrams) + + require.Error(t, serverConn.SendDatagram([]byte("foo"))) + require.NoError(t, clientConn.SendDatagram([]byte("bar"))) + datagram, err := serverConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("bar"), datagram) + }) +} + +func testDatagramNegotiation(t *testing.T, serverEnableDatagram, clientEnableDatagram bool) { + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{EnableDatagrams: serverEnableDatagram}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientConn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{EnableDatagrams: clientEnableDatagram}), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + serverState := serverConn.ConnectionState().SupportsDatagrams + clientState := clientConn.ConnectionState().SupportsDatagrams + require.Equal(t, serverEnableDatagram, serverState.Local, "server local datagram support") + require.Equal(t, clientEnableDatagram, serverState.Remote, "server view of client datagram support") + require.Equal(t, clientEnableDatagram, clientState.Local, "client local datagram support") + require.Equal(t, serverEnableDatagram, clientState.Remote, "client view of server datagram support") + + if clientEnableDatagram { + require.NoError(t, serverConn.SendDatagram([]byte("foo"))) + datagram, err := clientConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("foo"), datagram) + } else { + require.Error(t, serverConn.SendDatagram([]byte("foo"))) + } + + if serverEnableDatagram { + require.NoError(t, clientConn.SendDatagram([]byte("bar"))) + datagram, err := serverConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, []byte("bar"), datagram) + } else { + require.Error(t, clientConn.SendDatagram([]byte("bar"))) + } +} + +func TestDatagramSizeLimit(t *testing.T) { + const maxDatagramSize = 456 + originalMaxDatagramSize := wire.MaxDatagramSize + wire.MaxDatagramSize = maxDatagramSize + t.Cleanup(func() { wire.MaxDatagramSize = originalMaxDatagramSize }) + + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientConn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + err = clientConn.SendDatagram(bytes.Repeat([]byte("a"), maxDatagramSize+100)) // definitely too large + require.Error(t, err) + var sizeErr *quic.DatagramTooLargeError + require.ErrorAs(t, err, &sizeErr) + require.InDelta(t, sizeErr.MaxDatagramPayloadSize, maxDatagramSize, 10) + + require.NoError(t, clientConn.SendDatagram(bytes.Repeat([]byte("b"), int(sizeErr.MaxDatagramPayloadSize)))) + require.Error(t, clientConn.SendDatagram(bytes.Repeat([]byte("c"), int(sizeErr.MaxDatagramPayloadSize+1)))) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + datagram, err := serverConn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.Equal(t, bytes.Repeat([]byte("b"), int(sizeErr.MaxDatagramPayloadSize)), datagram) +} + +func TestDatagramSizeLimitWithMTUDiscovery(t *testing.T) { + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + var eventRecorder events.Recorder + clientConn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + InitialPacketSize: protocol.MinInitialPacketSize, + EnableDatagrams: true, + Tracer: newTracer(&eventRecorder), + }), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + serverErrChan := make(chan error, 1) + go func() { + str, err := serverConn.AcceptStream(ctx) + if err != nil { + serverErrChan <- err + return + } + _, err = io.Copy(io.Discard, str) + serverErrChan <- err + }() + + str, err := clientConn.OpenStream() + require.NoError(t, err) + + data := bytes.Repeat([]byte("d"), 16*1024) + var discoveredMTU int + for discoveredMTU == 0 { + _, err = str.Write(data) + require.NoError(t, err) + events := eventRecorder.Events(qlog.MTUUpdated{}) + if len(events) > 0 { + update := events[len(events)-1].(qlog.MTUUpdated) + if update.Done { + discoveredMTU = update.Value + } + } + require.NoError(t, ctx.Err()) + } + require.NoError(t, str.Close()) + + // Receiving the stream FIN guarantees that the client applied the MTU update observed above. + select { + case err := <-serverErrChan: + require.NoError(t, err) + case <-ctx.Done(): + require.NoError(t, ctx.Err()) + } + + err = clientConn.SendDatagram(bytes.Repeat([]byte("x"), 2000)) + var sizeErr *quic.DatagramTooLargeError + require.ErrorAs(t, err, &sizeErr) + maxPayloadSize := sizeErr.MaxDatagramPayloadSize + require.Greater(t, maxPayloadSize, int64(protocol.MinInitialPacketSize), "MTU discovery should increase the datagram size limit") + require.Less(t, maxPayloadSize, int64(discoveredMTU), "datagram payload must leave room for packet overhead") + + datagramData := bytes.Repeat([]byte("z"), int(maxPayloadSize)) + require.NoError(t, clientConn.SendDatagram(datagramData)) + datagram, err := serverConn.ReceiveDatagram(ctx) + require.NoError(t, err, "datagram should be deliverable when respecting MaxDatagramPayloadSize") + require.Equal(t, datagramData, datagram) +} + +func TestDatagramLoss(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 100 * time.Millisecond + const numDatagrams = 100 + const datagramSize = 500 + + clientAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001} + serverAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002} + var droppedToClient, droppedToServer, total atomic.Int32 + n := &simnet.Simnet{ + Router: &directionAwareDroppingRouter{ + ClientAddr: clientAddr, + ServerAddr: serverAddr, + Drop: func(d direction, p simnet.Packet) bool { + if wire.IsLongHeaderPacket(p.Data[0]) { // don't drop Long Header packets + return false + } + if len(p.Data) < datagramSize { // don't drop ACK-only packets + return false + } + total.Add(1) + // drop about 20% of Short Header packets with DATAGRAM frames + if mrand.Int()%5 == 0 { + switch d { + case directionToClient: + droppedToClient.Add(1) + case directionToServer: + droppedToServer.Add(1) + } + return true + } + return false + }, + }, + } + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + clientPacketConn := n.NewEndpoint(clientAddr, settings) + defer clientPacketConn.Close() + serverPacketConn := n.NewEndpoint(serverAddr, settings) + defer serverPacketConn.Close() + require.NoError(t, n.Start()) + defer n.Close() + + server, err := quic.Listen( + serverPacketConn, + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true, EnableDatagrams: true}), + ) + require.NoError(t, err) + defer server.Close() + + const sendInterval = time.Second // send a datagram every second + ctx, cancel := context.WithTimeout(context.Background(), (numDatagrams+10)*sendInterval) + defer cancel() + clientConn, err := quic.Dial( + ctx, + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true, EnableDatagrams: true}), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + var clientDatagrams, serverDatagrams int + clientErrChan := make(chan error, 1) + go func() { + defer close(clientErrChan) + for { + if _, err := clientConn.ReceiveDatagram(ctx); err != nil { + clientErrChan <- err + return + } + clientDatagrams++ + } + }() + + for i := range numDatagrams { + payload := bytes.Repeat([]byte{uint8(i)}, datagramSize) + require.NoError(t, clientConn.SendDatagram(payload)) + require.NoError(t, serverConn.SendDatagram(payload)) + time.Sleep(sendInterval) + } + + serverErrChan := make(chan error, 1) + go func() { + defer close(serverErrChan) + for { + if _, err := serverConn.ReceiveDatagram(ctx); err != nil { + serverErrChan <- err + return + } + serverDatagrams++ + } + }() + + select { + case err := <-clientErrChan: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(5 * numDatagrams * sendInterval): + t.Fatal("timeout") + } + select { + case err := <-serverErrChan: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(5 * numDatagrams * sendInterval): + t.Fatal("timeout") + } + + numDroppedToClient := droppedToClient.Load() + numDroppedToServer := droppedToServer.Load() + t.Logf("dropped %d to client and %d to server out of %d packets", numDroppedToClient, numDroppedToServer, total.Load()) + assert.NotZero(t, numDroppedToClient) + assert.NotZero(t, numDroppedToServer) + t.Logf("server received %d out of %d sent datagrams", serverDatagrams, numDatagrams) + assert.EqualValues(t, numDatagrams-numDroppedToServer, serverDatagrams, "datagrams received by the server") + t.Logf("client received %d out of %d sent datagrams", clientDatagrams, numDatagrams) + assert.EqualValues(t, numDatagrams-numDroppedToClient, clientDatagrams, "datagrams received by the client") + }) +} diff --git a/third_party/quic-go/integrationtests/self/deadline_test.go b/third_party/quic-go/integrationtests/self/deadline_test.go new file mode 100644 index 0000000..4804912 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/deadline_test.go @@ -0,0 +1,235 @@ +package self_test + +import ( + "bytes" + "context" + "fmt" + "io" + "net" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/require" +) + +func setupDeadlineTest(t *testing.T) (serverStr, clientStr *quic.Stream) { + t.Helper() + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + t.Cleanup(func() { server.Close() }) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + t.Cleanup(func() { conn.CloseWithError(0, "") }) + clientStr, err = conn.OpenStream() + require.NoError(t, err) + _, err = clientStr.Write([]byte{0}) // need to write one byte so the server learns about the stream + require.NoError(t, err) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + t.Cleanup(func() { serverConn.CloseWithError(0, "") }) + serverStr, err = serverConn.AcceptStream(ctx) + require.NoError(t, err) + + _, err = serverStr.Read([]byte{0}) + require.NoError(t, err) + return serverStr, clientStr +} + +func TestReadDeadlineSync(t *testing.T) { + serverStr, clientStr := setupDeadlineTest(t) + + const timeout = time.Millisecond + errChan := make(chan error, 1) + go func() { + _, err := serverStr.Write(PRDataLong) + errChan <- err + }() + + var bytesRead int + var timeoutCounter int + buf := make([]byte, 1<<10) + data := make([]byte, len(PRDataLong)) + clientStr.SetReadDeadline(time.Now().Add(timeout)) + for bytesRead < len(PRDataLong) { + n, err := clientStr.Read(buf) + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + timeoutCounter++ + clientStr.SetReadDeadline(time.Now().Add(timeout)) + } else { + require.NoError(t, err) + } + copy(data[bytesRead:], buf[:n]) + bytesRead += n + } + require.Equal(t, PRDataLong, data) + // make sure the test actually worked and Read actually ran into the deadline a few times + t.Logf("ran into deadline %d times", timeoutCounter) + require.GreaterOrEqual(t, timeoutCounter, 10) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestReadDeadlineAsync(t *testing.T) { + serverStr, clientStr := setupDeadlineTest(t) + + const timeout = time.Millisecond + errChan := make(chan error, 1) + go func() { + _, err := serverStr.Write(PRDataLong) + errChan <- err + }() + + var bytesRead int + var timeoutCounter int + buf := make([]byte, 1<<10) + data := make([]byte, len(PRDataLong)) + received := make(chan struct{}) + go func() { + for { + select { + case <-received: + return + default: + time.Sleep(timeout) + } + clientStr.SetReadDeadline(time.Now().Add(timeout)) + } + }() + + for bytesRead < len(PRDataLong) { + n, err := clientStr.Read(buf) + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + timeoutCounter++ + } else { + require.NoError(t, err) + } + copy(data[bytesRead:], buf[:n]) + bytesRead += n + } + + require.Equal(t, PRDataLong, data) + close(received) + + // make sure the test actually worked and Read actually ran into the deadline a few times + t.Logf("ran into deadline %d times", timeoutCounter) + require.GreaterOrEqual(t, timeoutCounter, 10) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestWriteDeadlineSync(t *testing.T) { + serverStr, clientStr := setupDeadlineTest(t) + + const timeout = time.Millisecond + + errChan := make(chan error, 1) + go func() { + defer close(errChan) + data, err := io.ReadAll(serverStr) + if err != nil { + errChan <- err + } + if !bytes.Equal(PRDataLong, data) { + errChan <- fmt.Errorf("data mismatch") + } + }() + + var bytesWritten int + var timeoutCounter int + clientStr.SetWriteDeadline(time.Now().Add(timeout)) + for bytesWritten < len(PRDataLong) { + n, err := clientStr.Write(PRDataLong[bytesWritten:]) + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + timeoutCounter++ + clientStr.SetWriteDeadline(time.Now().Add(timeout)) + } else { + require.NoError(t, err) + } + bytesWritten += n + } + clientStr.Close() + + // make sure the test actually worked and Write actually ran into the deadline a few times + t.Logf("ran into deadline %d times", timeoutCounter) + require.GreaterOrEqual(t, timeoutCounter, 10) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestWriteDeadlineAsync(t *testing.T) { + serverStr, clientStr := setupDeadlineTest(t) + + const timeout = time.Millisecond + + errChan := make(chan error, 1) + go func() { + defer close(errChan) + data, err := io.ReadAll(serverStr) + if err != nil { + errChan <- err + } + if !bytes.Equal(PRDataLong, data) { + errChan <- fmt.Errorf("data mismatch") + } + }() + + clientStr.SetWriteDeadline(time.Now().Add(timeout)) + readDone := make(chan struct{}) + deadlineDone := make(chan struct{}) + go func() { + defer close(deadlineDone) + for { + select { + case <-readDone: + return + default: + time.Sleep(timeout) + } + clientStr.SetWriteDeadline(time.Now().Add(timeout)) + } + }() + + var bytesWritten int + var timeoutCounter int + clientStr.SetWriteDeadline(time.Now().Add(timeout)) + for bytesWritten < len(PRDataLong) { + n, err := clientStr.Write(PRDataLong[bytesWritten:]) + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + timeoutCounter++ + } else { + require.NoError(t, err) + } + bytesWritten += n + } + clientStr.Close() + + close(readDone) + + // make sure the test actually worked and Write actually ran into the deadline a few times + t.Logf("ran into deadline %d times", timeoutCounter) + require.GreaterOrEqual(t, timeoutCounter, 10) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} diff --git a/third_party/quic-go/integrationtests/self/drop_test.go b/third_party/quic-go/integrationtests/self/drop_test.go new file mode 100644 index 0000000..fc16fa3 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/drop_test.go @@ -0,0 +1,117 @@ +package self_test + +import ( + "context" + "fmt" + "net" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/require" +) + +func TestPacketDrops(t *testing.T) { + for _, direction := range []protocol.Perspective{protocol.PerspectiveClient, protocol.PerspectiveServer} { + t.Run(fmt.Sprintf("from %s", direction), func(t *testing.T) { + testPacketDrops(t, direction) + }) + } +} + +func testPacketDrops(t *testing.T, direction protocol.Perspective) { + synctest.Test(t, func(t *testing.T) { + const numMessages = 50 + const rtt = 10 * time.Millisecond + + addrClient := &net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001} + addrServer := &net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002} + + var numDroppedPackets atomic.Int32 + messageInterval := randomDuration(10*time.Millisecond, 100*time.Millisecond) + dropDuration := randomDuration(messageInterval*3, 2*time.Second) + dropDelay := randomDuration(25*time.Millisecond, numMessages*messageInterval/2) + + startTime := time.Now() + n := &simnet.Simnet{ + Router: &droppingRouter{ + Drop: func(p simnet.Packet) bool { + switch p.To { + case addrClient: + if direction == protocol.PerspectiveClient { + return false + } + case addrServer: + if direction == protocol.PerspectiveServer { + return false + } + } + if wire.IsLongHeaderPacket(p.Data[0]) { // don't interfere with the handshake + return false + } + drop := time.Now().After(startTime.Add(dropDelay)) && time.Now().Before(startTime.Add(dropDelay).Add(dropDuration)) + if drop { + numDroppedPackets.Add(1) + } + return drop + }, + }, + } + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + clientPacketConn := n.NewEndpoint(addrClient, settings) + defer clientPacketConn.Close() + serverPacketConn := n.NewEndpoint(addrServer, settings) + defer serverPacketConn.Close() + + require.NoError(t, n.Start()) + defer n.Close() + + t.Logf("sending a message every %s, %d times", messageInterval, numMessages) + t.Logf("dropping packets for %s, after a delay of %s", dropDuration, dropDelay) + + ln, err := quic.Listen(serverPacketConn, getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + conn, err := quic.Dial(ctx, clientPacketConn, ln.Addr().(*net.UDPAddr), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + serverStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { + for i := range numMessages { + time.Sleep(messageInterval) + if _, err := serverStr.Write([]byte{uint8(i + 1)}); err != nil { + errChan <- err + return + } + } + }() + + str, err := conn.AcceptUniStream(ctx) + require.NoError(t, err) + for i := range numMessages { + b := []byte{0} + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 1, n) + require.Equal(t, byte(i+1), b[0]) + } + numDropped := numDroppedPackets.Load() + t.Logf("dropped %d packets", numDropped) + require.NotZero(t, numDropped) + }) +} diff --git a/third_party/quic-go/integrationtests/self/early_data_test.go b/third_party/quic-go/integrationtests/self/early_data_test.go new file mode 100644 index 0000000..3ca8e3c --- /dev/null +++ b/third_party/quic-go/integrationtests/self/early_data_test.go @@ -0,0 +1,72 @@ +package self_test + +import ( + "context" + "io" + "net" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + + "github.com/stretchr/testify/require" +) + +func TestEarlyData(t *testing.T) { + const rtt = 80 * time.Millisecond + ln, err := quic.ListenEarly(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + proxy := &quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { return rtt / 2 }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + connChan := make(chan *quic.Conn) + errChan := make(chan error) + go func() { + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan <- err + return + } + connChan <- conn + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + clientConn, err := quic.Dial(ctx, newUDPConnLocalhost(t), proxy.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + var serverConn *quic.Conn + select { + case serverConn = <-connChan: + case err := <-errChan: + t.Fatalf("error accepting connection: %s", err) + } + str, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write([]byte("early data")) + require.NoError(t, err) + require.NoError(t, str.Close()) + // the write should have completed before the handshake + select { + case <-serverConn.HandshakeComplete(): + t.Fatal("handshake shouldn't be completed yet") + default: + } + + clientStr, err := clientConn.AcceptUniStream(context.Background()) + require.NoError(t, err) + data, err := io.ReadAll(clientStr) + require.NoError(t, err) + require.Equal(t, []byte("early data"), data) + + clientConn.CloseWithError(0, "") + <-serverConn.Context().Done() +} diff --git a/third_party/quic-go/integrationtests/self/handshake_context_test.go b/third_party/quic-go/integrationtests/self/handshake_context_test.go new file mode 100644 index 0000000..d619e0c --- /dev/null +++ b/third_party/quic-go/integrationtests/self/handshake_context_test.go @@ -0,0 +1,289 @@ +package self_test + +import ( + "context" + "crypto/tls" + "errors" + "net" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/stretchr/testify/require" +) + +func TestHandshakeContextTimeout(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(20*time.Millisecond)) + defer cancel() + + conn := newUDPConnLocalhost(t) + + errChan := make(chan error, 1) + go func() { + _, err := quic.Dial(ctx, newUDPConnLocalhost(t), conn.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + errChan <- err + }() + + require.ErrorIs(t, <-errChan, context.DeadlineExceeded) +} + +func TestHandshakeCancellationError(t *testing.T) { + ctx, cancel := context.WithCancelCause(context.Background()) + errChan := make(chan error, 1) + conn := newUDPConnLocalhost(t) + go func() { + _, err := quic.Dial(ctx, newUDPConnLocalhost(t), conn.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + errChan <- err + }() + + cancel(errors.New("application cancelled")) + require.EqualError(t, <-errChan, "application cancelled") +} + +func TestConnContextOnServerSide(t *testing.T) { + tlsGetConfigForClientContextChan := make(chan context.Context, 1) + tlsGetCertificateContextChan := make(chan context.Context, 1) + tracerContextChan := make(chan context.Context, 1) + connContextChan := make(chan context.Context, 1) + streamContextChan := make(chan context.Context, 1) + + tr := &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnContext: func(ctx context.Context, _ *quic.ClientInfo) (context.Context, error) { + return context.WithValue(ctx, "foo", "bar"), nil + }, + } + defer tr.Close() + + server, err := tr.Listen( + &tls.Config{ + GetConfigForClient: func(info *tls.ClientHelloInfo) (*tls.Config, error) { + tlsGetConfigForClientContextChan <- info.Context() + tlsConf := getTLSConfig() + tlsConf.GetCertificate = func(info *tls.ClientHelloInfo) (*tls.Certificate, error) { + tlsGetCertificateContextChan <- info.Context() + return &tlsConf.Certificates[0], nil + } + return tlsConf, nil + }, + }, + getQuicConfig(&quic.Config{ + Tracer: func(ctx context.Context, _ bool, _ quic.ConnectionID) qlogwriter.Trace { + tracerContextChan <- ctx + return nil + }, + }), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + c, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + connContextChan <- serverConn.Context() + str, err := serverConn.OpenUniStream() + require.NoError(t, err) + streamContextChan <- str.Context() + str.Write([]byte{1, 2, 3}) + + _, err = c.AcceptUniStream(ctx) + require.NoError(t, err) + c.CloseWithError(1337, "bye") + + checkContext := func(c <-chan context.Context, checkCancellationCause bool) { + t.Helper() + var ctx context.Context + select { + case ctx = <-c: + case <-time.After(time.Second): + t.Fatal("timeout waiting for context") + } + + val := ctx.Value("foo") + require.NotNil(t, val) + v := val.(string) + require.Equal(t, "bar", v) + + select { + case <-ctx.Done(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for context to be done") + } + + if !checkCancellationCause { + return + } + ctxErr := context.Cause(ctx) + var appErr *quic.ApplicationError + require.ErrorAs(t, ctxErr, &appErr) + require.Equal(t, quic.ApplicationErrorCode(1337), appErr.ErrorCode) + } + + checkContext(connContextChan, true) + checkContext(tracerContextChan, true) + checkContext(streamContextChan, true) + // crypto/tls cancels the context when the TLS handshake completes. + checkContext(tlsGetConfigForClientContextChan, false) + checkContext(tlsGetCertificateContextChan, false) +} + +func TestConnContextRejection(t *testing.T) { + t.Run("rejecting", func(t *testing.T) { + testConnContextRejection(t, true) + }) + t.Run("not rejecting", func(t *testing.T) { + testConnContextRejection(t, false) + }) +} + +func testConnContextRejection(t *testing.T, reject bool) { + tr := &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnContext: func(ctx context.Context, ci *quic.ClientInfo) (context.Context, error) { + if reject { + return nil, errors.New("rejecting connection") + } + return context.WithValue(ctx, "addr", ci.RemoteAddr), nil + }, + } + defer tr.Close() + + server, err := tr.Listen( + getTLSConfig(), + getQuicConfig(nil), + ) + require.NoError(t, err) + defer server.Close() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + pc := newUDPConnLocalhost(t) + c, err := quic.Dial(ctx, pc, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + if reject { + require.ErrorIs(t, err, &quic.TransportError{Remote: true, ErrorCode: quic.ConnectionRefused}) + return + } + require.NoError(t, err) + defer c.CloseWithError(0, "") + + conn, err := server.Accept(ctx) + require.NoError(t, err) + require.Equal(t, pc.LocalAddr().String(), conn.Context().Value("addr").(net.Addr).String()) + conn.CloseWithError(0, "") +} + +// Users are not supposed to return a fresh context from ConnContext, but we should handle it gracefully. +func TestConnContextFreshContext(t *testing.T) { + tr := &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnContext: func(ctx context.Context, _ *quic.ClientInfo) (context.Context, error) { + return context.Background(), nil + }, + } + defer tr.Close() + server, err := tr.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + errChan := make(chan error, 1) + go func() { + conn, err := server.Accept(context.Background()) + if err != nil { + errChan <- err + return + } + conn.CloseWithError(1337, "bye") + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + c, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + select { + case <-c.Context().Done(): + case err := <-errChan: + t.Fatalf("accept failed: %v", err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestContextOnClientSide(t *testing.T) { + tlsServerConf := getTLSConfig() + tlsServerConf.ClientAuth = tls.RequestClientCert + server, err := quic.Listen(newUDPConnLocalhost(t), tlsServerConf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + tlsContextChan := make(chan context.Context, 1) + tracerContextChan := make(chan context.Context, 1) + tlsConf := getTLSClientConfig() + tlsConf.GetClientCertificate = func(info *tls.CertificateRequestInfo) (*tls.Certificate, error) { + tlsContextChan <- info.Context() + return &tlsServerConf.Certificates[0], nil + } + + ctx, cancel := context.WithCancel(context.WithValue(context.Background(), "foo", "bar")) + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + tlsConf, + getQuicConfig(&quic.Config{ + Tracer: func(ctx context.Context, _ bool, _ quic.ConnectionID) qlogwriter.Trace { + tracerContextChan <- ctx + return nil + }, + }), + ) + require.NoError(t, err) + cancel() + + // Make sure the connection context is not cancelled (even though derived from the ctx passed to Dial) + select { + case <-conn.Context().Done(): + t.Fatal("context should not be cancelled") + default: + } + + checkContext := func(ctx context.Context, checkCancellationCause bool) { + t.Helper() + val := ctx.Value("foo") + require.NotNil(t, val) + require.Equal(t, "bar", val.(string)) + if !checkCancellationCause { + return + } + ctxErr := context.Cause(ctx) + var appErr *quic.ApplicationError + require.ErrorAs(t, ctxErr, &appErr) + require.EqualValues(t, 1337, appErr.ErrorCode) + } + + checkContextFromChan := func(c <-chan context.Context, checkCancellationCause bool) { + t.Helper() + var ctx context.Context + select { + case ctx = <-c: + case <-time.After(time.Second): + t.Fatal("timeout waiting for context") + } + checkContext(ctx, checkCancellationCause) + } + + str, err := conn.OpenUniStream() + require.NoError(t, err) + conn.CloseWithError(1337, "bye") + + checkContext(conn.Context(), true) + checkContext(str.Context(), true) + // crypto/tls cancels the context when the TLS handshake completes + checkContextFromChan(tlsContextChan, false) + checkContextFromChan(tracerContextChan, false) +} diff --git a/third_party/quic-go/integrationtests/self/handshake_drop_test.go b/third_party/quic-go/integrationtests/self/handshake_drop_test.go new file mode 100644 index 0000000..618c848 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/handshake_drop_test.go @@ -0,0 +1,411 @@ +package self_test + +import ( + "bytes" + "context" + "crypto/tls" + "fmt" + "io" + mrand "math/rand/v2" + "net" + "slices" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/testutils/events" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/require" +) + +func dropTestProtocolClientSpeaksFirst(t *testing.T, ln *quic.Listener, clientConn net.PacketConn, clientConf *tls.Config, timeout time.Duration, data []byte) *quic.Conn { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + conn, err := quic.Dial( + ctx, + clientConn, + ln.Addr(), + clientConf, + getQuicConfig(&quic.Config{ + MaxIdleTimeout: timeout, + HandshakeIdleTimeout: timeout, + DisablePathMTUDiscovery: true, + }), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + str, err := conn.OpenUniStream() + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { + defer str.Close() + _, err := str.Write(data) + errChan <- err + }() + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + serverStr, err := serverConn.AcceptUniStream(ctx) + require.NoError(t, err) + b, err := io.ReadAll(&readerWithTimeout{Reader: serverStr, Timeout: timeout}) + require.NoError(t, err) + require.Equal(t, b, data) + serverConn.CloseWithError(0, "") + + return conn +} + +func dropTestProtocolServerSpeaksFirst(t *testing.T, ln *quic.Listener, clientConn net.PacketConn, clientConf *tls.Config, timeout time.Duration, data []byte) *quic.Conn { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + conn, err := quic.Dial( + ctx, + clientConn, + ln.Addr(), + clientConf, + getQuicConfig(&quic.Config{ + MaxIdleTimeout: timeout, + HandshakeIdleTimeout: timeout, + DisablePathMTUDiscovery: true, + }), + ) + require.NoError(t, err) + + errChan := make(chan error, 1) + go func() { + defer close(errChan) + defer conn.CloseWithError(0, "") + str, err := conn.AcceptUniStream(ctx) + if err != nil { + errChan <- err + return + } + b, err := io.ReadAll(&readerWithTimeout{Reader: str, Timeout: timeout}) + if err != nil { + errChan <- err + return + } + if !bytes.Equal(b, data) { + errChan <- fmt.Errorf("data mismatch: %x != %x", b, data) + return + } + }() + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + serverStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + _, err = serverStr.Write(data) + require.NoError(t, err) + require.NoError(t, serverStr.Close()) + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(timeout): + t.Fatal("server connection not closed") + } + + select { + case <-conn.Context().Done(): + case <-time.After(timeout): + t.Fatal("server connection not closed") + } + + return conn +} + +func dropTestProtocolNobodySpeaks(t *testing.T, ln *quic.Listener, clientConn net.PacketConn, clientConf *tls.Config, timeout time.Duration, _ []byte) *quic.Conn { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + conn, err := quic.Dial( + ctx, + clientConn, + ln.Addr(), + clientConf, + getQuicConfig(&quic.Config{ + MaxIdleTimeout: timeout, + HandshakeIdleTimeout: timeout, + DisablePathMTUDiscovery: true, + }), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + serverConn.CloseWithError(0, "") + + return conn +} + +func dropCallbackDropNthPacket(dir direction, ns ...int) func(direction, simnet.Packet) bool { + var toClient, toServer atomic.Int32 + return func(d direction, p simnet.Packet) bool { + switch d { + case directionToClient: + c := toClient.Add(1) + if d == dir || dir == directionBoth { + return slices.Contains(ns, int(c)) + } + case directionToServer: + c := toServer.Add(1) + if dir == d || dir == directionBoth { + return slices.Contains(ns, int(c)) + } + } + return false + } +} + +func dropCallbackDropOneThird(_ direction) func(direction, simnet.Packet) bool { + const maxSequentiallyDropped = 10 + var mx sync.Mutex + var toClient, toServer int + return func(d direction, p simnet.Packet) bool { + drop := mrand.IntN(3) == 0 + + mx.Lock() + defer mx.Unlock() + // never drop more than 10 consecutive packets + if d == directionToClient || d == directionBoth { + if drop { + toClient++ + if toClient > maxSequentiallyDropped { + drop = false + } + } + if !drop { + toClient = 0 + } + } + if d == directionToServer || d == directionBoth { + if drop { + toServer++ + if toServer > maxSequentiallyDropped { + drop = false + } + } + if !drop { + toServer = 0 + } + } + return drop + } +} + +func TestHandshakeWithPacketLoss(t *testing.T) { + data := GeneratePRData(5000) + const timeout = 2 * time.Minute + const rtt = 20 * time.Millisecond + + type dropPattern string + + const ( + dropPatternDrop1stPacket dropPattern = "drop 1st packet" + dropPatternDropFirst3Packets dropPattern = "drop first 3 packets" + dropPatternDropOneThirdOfPackets dropPattern = "drop 1/3 of packets" + ) + + type testConfig struct { + postQuantum bool + longCertChain bool + doRetry bool + } + + for _, dir := range []direction{directionToClient, directionToServer, directionBoth} { + for _, pattern := range []dropPattern{ + dropPatternDrop1stPacket, + dropPatternDropFirst3Packets, + dropPatternDropOneThirdOfPackets, + } { + t.Run(fmt.Sprintf("%s in direction %s", pattern, dir), func(t *testing.T) { + for _, conf := range []testConfig{ + {postQuantum: false, longCertChain: false, doRetry: true}, + {postQuantum: false, longCertChain: false, doRetry: false}, + {postQuantum: false, longCertChain: true, doRetry: false}, + {postQuantum: true, longCertChain: false, doRetry: false}, + {postQuantum: true, longCertChain: true, doRetry: false}, + } { + for _, test := range []struct { + name string + fn func(t *testing.T, ln *quic.Listener, clientConn net.PacketConn, clientConf *tls.Config, timeout time.Duration, data []byte) *quic.Conn + }{ + {"client speaks first", dropTestProtocolClientSpeaksFirst}, + {"server speaks first", dropTestProtocolServerSpeaksFirst}, + {"nobody speaks", dropTestProtocolNobodySpeaks}, + } { + t.Run(fmt.Sprintf("retry: %t/%s", conf.doRetry, test.name), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + clientAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001} + serverAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002} + var fn func(direction, simnet.Packet) bool + switch pattern { + case dropPatternDrop1stPacket: + fn = dropCallbackDropNthPacket(dir, 1) + case dropPatternDropFirst3Packets: + fn = dropCallbackDropNthPacket(dir, 1, 2, 3) + case dropPatternDropOneThirdOfPackets: + fn = dropCallbackDropOneThird(dir) + } + var numDropped atomic.Int32 + n := &simnet.Simnet{ + Router: &directionAwareDroppingRouter{ + ClientAddr: clientAddr, + ServerAddr: serverAddr, + Drop: func(d direction, p simnet.Packet) bool { + drop := fn(d, p) + if drop { + numDropped.Add(1) + } + return drop + }, + }, + } + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + clientConn := n.NewEndpoint(clientAddr, settings) + defer clientConn.Close() + serverConn := n.NewEndpoint(serverAddr, settings) + defer serverConn.Close() + require.NoError(t, n.Start()) + defer n.Close() + + var tlsConf *tls.Config + if conf.longCertChain { + tlsConf = getTLSConfigWithLongCertChain() + } else { + tlsConf = getTLSConfig() + } + clientConf := getTLSClientConfig() + if !conf.postQuantum { + clientConf.CurvePreferences = []tls.CurveID{tls.CurveP384} + } + + tr := &quic.Transport{ + Conn: serverConn, + VerifySourceAddress: func(net.Addr) bool { return conf.doRetry }, + } + defer tr.Close() + + ln, err := tr.Listen( + tlsConf, + getQuicConfig(&quic.Config{ + MaxIdleTimeout: timeout, + HandshakeIdleTimeout: timeout, + DisablePathMTUDiscovery: true, + }), + ) + require.NoError(t, err) + defer ln.Close() + + conn := test.fn(t, ln, clientConn, clientConf, timeout, data) + curveID := getCurveID(conn.ConnectionState().TLS) + if conf.postQuantum { + require.Equal(t, tls.X25519MLKEM768, curveID) + } else { + require.Equal(t, tls.CurveP384, curveID) + } + + if pattern != dropPatternDropOneThirdOfPackets { + require.NotZero(t, numDropped.Load()) + } + t.Logf("dropped %d packets", numDropped.Load()) + }) + }) + } + } + }) + } + } +} + +func TestHandshakePacketBuffering(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 20 * time.Millisecond + + clientAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001} + serverAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002} + var droppedFirst atomic.Bool + n := &simnet.Simnet{ + Router: &directionAwareDroppingRouter{ + ClientAddr: clientAddr, + ServerAddr: serverAddr, + Drop: func(d direction, p simnet.Packet) bool { + if droppedFirst.Load() { + return false + } + if d == directionToClient && containsPacketType(p.Data, protocol.PacketTypeInitial) { + droppedFirst.Store(true) + return true + } + return false + }, + }, + } + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + clientConn := n.NewEndpoint(clientAddr, settings) + defer clientConn.Close() + serverConn := n.NewEndpoint(serverAddr, settings) + defer serverConn.Close() + require.NoError(t, n.Start()) + defer n.Close() + + var serverEventRecorder events.Recorder + ln, err := quic.Listen( + serverConn, + getTLSConfig(), + getQuicConfig(&quic.Config{Tracer: newTracer(&serverEventRecorder)}), + ) + require.NoError(t, err) + defer ln.Close() + + var clientEventRecorder events.Recorder + conn, err := quic.Dial( + context.Background(), + clientConn, + ln.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{Tracer: newTracer(&clientEventRecorder)}), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + str, err := conn.OpenUniStream() + require.NoError(t, err) + data := []byte("foobar") + _, err = str.Write(data) + require.NoError(t, err) + require.NoError(t, str.Close()) + + require.Empty(t, serverEventRecorder.Events(qlog.PacketBuffered{})) + buffered := clientEventRecorder.Events(qlog.PacketBuffered{}) + t.Logf("buffered packets: %d", len(buffered)) + require.NotEmpty(t, buffered) + receivedPackets := make(map[qlog.DatagramPayloadChecksum][]qlog.PacketType) + for _, ev := range clientEventRecorder.Events(qlog.PacketReceived{}) { + checksum := ev.(qlog.PacketReceived).DatagramPayloadChecksum + receivedPackets[checksum] = append(receivedPackets[checksum], ev.(qlog.PacketReceived).Header.PacketType) + } + for _, ev := range buffered { + checksum := ev.(qlog.PacketBuffered).DatagramPayloadChecksum + require.Contains(t, receivedPackets, checksum) + require.Contains(t, receivedPackets[checksum], qlog.PacketTypeHandshake) + } + + sconn, err := ln.Accept(context.Background()) + require.NoError(t, err) + defer sconn.CloseWithError(0, "") + sstr, err := sconn.AcceptUniStream(context.Background()) + require.NoError(t, err) + b, err := io.ReadAll(sstr) + require.NoError(t, err) + require.Equal(t, data, b) + require.Equal(t, rtt, sconn.ConnectionStats().SmoothedRTT) + }) +} diff --git a/third_party/quic-go/integrationtests/self/handshake_rtt_test.go b/third_party/quic-go/integrationtests/self/handshake_rtt_test.go new file mode 100644 index 0000000..53ebc3c --- /dev/null +++ b/third_party/quic-go/integrationtests/self/handshake_rtt_test.go @@ -0,0 +1,201 @@ +package self_test + +import ( + "context" + "crypto/tls" + "io" + "net" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/require" +) + +func TestHandshakeRTTRetry(t *testing.T) { + t.Run("retry", func(t *testing.T) { + testHandshakeRTTRetry(t, true) + }) + t.Run("no retry", func(t *testing.T) { + testHandshakeRTTRetry(t, false) + }) +} + +func testHandshakeRTTRetry(t *testing.T, doRetry bool) { + var addrVerified bool + rtts := testHandshakeMeasureHandshake(t, + func(net.Addr) bool { return doRetry }, + getTLSConfig(), + getQuicConfig(&quic.Config{ + GetConfigForClient: func(info *quic.ClientInfo) (*quic.Config, error) { + addrVerified = info.AddrVerified + return nil, nil + }, + }), + ) + if doRetry { + require.True(t, addrVerified, "should have verified address") + require.GreaterOrEqual(t, rtts, float64(2)) + require.Less(t, rtts, float64(2.1)) + } else { + require.False(t, addrVerified, "should not have verified address") + require.GreaterOrEqual(t, rtts, float64(1)) + require.Less(t, rtts, float64(1.1)) + } +} + +func TestHandshakeRTTHelloRetryRequest(t *testing.T) { + tlsConf := getTLSConfig() + tlsConf.CurvePreferences = []tls.CurveID{tls.CurveP384} + rtts := testHandshakeMeasureHandshake(t, nil, tlsConf, getQuicConfig(nil)) + require.GreaterOrEqual(t, rtts, float64(2)) + require.Less(t, rtts, float64(2.1)) +} + +func testHandshakeMeasureHandshake(t *testing.T, verifySourceAddress func(net.Addr) bool, tlsConf *tls.Config, quicConf *quic.Config) float64 { + var rtts float64 + synctest.Test(t, func(t *testing.T) { + const rtt = 100 * time.Millisecond + + clientPacketConn, serverPacketConn, close := newSimnetLink(t, rtt) + defer close(t) + + tr := &quic.Transport{ + Conn: serverPacketConn, + VerifySourceAddress: verifySourceAddress, + } + addTracer(tr) + defer tr.Close() + ln, err := tr.Listen(tlsConf, quicConf) + require.NoError(t, err) + defer ln.Close() + + clientConfig := getQuicConfig(nil) + start := time.Now() + ctx, cancel := context.WithTimeout(context.Background(), 10*rtt) + defer cancel() + conn, err := quic.Dial( + ctx, + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + clientConfig, + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + rtts = time.Since(start).Seconds() / rtt.Seconds() + }) + return rtts +} + +func TestHandshake05RTT(t *testing.T) { + t.Run("using ListenEarly", func(t *testing.T) { + testHandshake05RTT(t, true) + }) + t.Run("using Listen", func(t *testing.T) { + testHandshake05RTT(t, false) + }) +} + +func testHandshake05RTT(t *testing.T, use05RTT bool) { + synctest.Test(t, func(t *testing.T) { + type accepter interface { + Accept(context.Context) (*quic.Conn, error) + } + + const rtt = 100 * time.Millisecond + clientPacketConn, serverPacketConn, close := newSimnetLink(t, rtt) + defer close(t) + var ln accepter + if use05RTT { + var err error + server, err := quic.ListenEarly(serverPacketConn, getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + ln = server + } else { + var err error + server, err := quic.Listen(serverPacketConn, getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + ln = server + } + + connChan := make(chan *quic.Conn, 1) + errChan := make(chan error, 1) + go func() { + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan <- err + return + } + str, err := conn.OpenUniStream() + if err != nil { + errChan <- err + return + } + if _, err := str.Write([]byte("foobar")); err != nil { + errChan <- err + return + } + if err := str.Close(); err != nil { + errChan <- err + return + } + + connChan <- conn + }() + + start := time.Now() + ctx, cancel := context.WithTimeout(context.Background(), 10*rtt) + defer cancel() + conn, err := quic.Dial( + ctx, + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(nil), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + rtts := time.Since(start).Seconds() / rtt.Seconds() + require.GreaterOrEqual(t, rtts, float64(1)) + require.Less(t, rtts, float64(1.1)) + + start = time.Now() + + select { + case err := <-errChan: + t.Fatal("failed to accept connection:", err) + case conn := <-connChan: + if !use05RTT { + // the server finishes the handshake 0.5 RTTs later + rtts = time.Since(start).Seconds() / rtt.Seconds() + require.GreaterOrEqual(t, rtts, float64(0.5)) + require.Less(t, rtts, float64(0.6)) + } + defer conn.CloseWithError(0, "") + } + + // If 0.5 RTT was used, the message should be received immediately, + // otherwise it should take 1 RTT. + str, err := conn.AcceptUniStream(ctx) + require.NoError(t, err) + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), data) + + rtts = time.Since(start).Seconds() / rtt.Seconds() + if use05RTT { + require.GreaterOrEqual(t, rtts, float64(0)) + require.Less(t, rtts, float64(0.1)) + } else { + require.GreaterOrEqual(t, rtts, float64(1)) + require.Less(t, rtts, float64(1.1)) + } + }) +} diff --git a/third_party/quic-go/integrationtests/self/handshake_test.go b/third_party/quic-go/integrationtests/self/handshake_test.go new file mode 100644 index 0000000..5dbe734 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/handshake_test.go @@ -0,0 +1,820 @@ +package self_test + +import ( + "context" + "crypto/fips140" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "runtime" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/qtls" + + "github.com/stretchr/testify/require" +) + +type tokenStore struct { + store quic.TokenStore + gets chan<- string + puts chan<- string +} + +var _ quic.TokenStore = &tokenStore{} + +func newTokenStore(gets, puts chan<- string) quic.TokenStore { + return &tokenStore{ + store: quic.NewLRUTokenStore(10, 4), + gets: gets, + puts: puts, + } +} + +func (c *tokenStore) Put(key string, token *quic.ClientToken) { + c.puts <- key + c.store.Put(key, token) +} + +func (c *tokenStore) Pop(key string) *quic.ClientToken { + c.gets <- key + return c.store.Pop(key) +} + +func TestHandshakeAddrResolutionHelpers(t *testing.T) { + server, err := quic.ListenAddr("localhost:0", getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.DialAddr( + ctx, + fmt.Sprintf("localhost:%d", server.Addr().(*net.UDPAddr).Port), + getTLSClientConfig(), + getQuicConfig(nil), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") +} + +func TestHandshake(t *testing.T) { + for _, tt := range []struct { + name string + conf *tls.Config + }{ + {"short cert chain", getTLSConfig()}, + {"long cert chain", getTLSConfigWithLongCertChain()}, + } { + t.Run(tt.name, func(t *testing.T) { + server, err := quic.Listen(newUDPConnLocalhost(t), tt.conf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + }) + } +} + +func TestHandshakeServerMismatch(t *testing.T) { + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + conf := getTLSClientConfig() + conf.ServerName = "foo.bar" + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err = quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), conf, getQuicConfig(nil)) + require.Error(t, err) + var transportErr *quic.TransportError + require.True(t, errors.As(err, &transportErr)) + require.True(t, transportErr.ErrorCode.IsCryptoError()) + require.Contains(t, transportErr.Error(), "x509: certificate is valid for localhost, not foo.bar") + var certErr *tls.CertificateVerificationError + require.True(t, errors.As(transportErr, &certErr)) +} + +func TestHandshakeCipherSuites(t *testing.T) { + for _, suiteID := range []uint16{ + tls.TLS_AES_128_GCM_SHA256, + tls.TLS_AES_256_GCM_SHA384, + tls.TLS_CHACHA20_POLY1305_SHA256, + } { + t.Run(tls.CipherSuiteName(suiteID), func(t *testing.T) { + if fips140.Enabled() && suiteID == tls.TLS_CHACHA20_POLY1305_SHA256 { + t.Skip("ChaCha20-Poly1305 is not allowed in FIPS 140-3 mode") + } + + reset := qtls.SetCipherSuite(suiteID) + defer reset() + + ln, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := ln.Accept(context.Background()) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + serverStr, err := serverConn.OpenStream() + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { + defer serverStr.Close() + _, err = serverStr.Write(PRData) + errChan <- err + }() + require.NoError(t, <-errChan) + + str, err := conn.AcceptStream(context.Background()) + require.NoError(t, err) + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRData, data) + require.Equal(t, suiteID, conn.ConnectionState().TLS.CipherSuite) + }) + } +} + +func TestTLSGetConfigForClientError(t *testing.T) { + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr) + defer tr.Close() + + tlsConf := &tls.Config{ + GetConfigForClient: func(info *tls.ClientHelloInfo) (*tls.Config, error) { + return nil, errors.New("nope") + }, + } + ln, err := tr.Listen(tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err = quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.True(t, transportErr.ErrorCode.IsCryptoError()) +} + +// Since we're not operating on a net.Conn, we need to jump through some hoops to set the addresses on the tls.ClientHelloInfo. +// Use a recursive setup to test that this works under all conditions. +func TestTLSConfigGetConfigForClientAddresses(t *testing.T) { + var local, remote net.Addr + var local2, remote2 net.Addr + done := make(chan struct{}) + tlsConf := &tls.Config{ + GetConfigForClient: func(info *tls.ClientHelloInfo) (*tls.Config, error) { + local = info.Conn.LocalAddr() + remote = info.Conn.RemoteAddr() + conf := getTLSConfig() + conf.GetCertificate = func(info *tls.ClientHelloInfo) (*tls.Certificate, error) { + defer close(done) + local2 = info.Conn.LocalAddr() + remote2 = info.Conn.RemoteAddr() + return &(conf.Certificates[0]), nil + } + return conf, nil + }, + } + server, err := quic.Listen(newUDPConnLocalhost(t), tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout waiting for GetCertificate callback") + } + + require.Equal(t, server.Addr(), local) + require.Equal(t, conn.LocalAddr().(*net.UDPAddr).Port, remote.(*net.UDPAddr).Port) + require.Equal(t, local, local2) + require.Equal(t, remote, remote2) +} + +func TestHandshakeFailsWithoutClientCert(t *testing.T) { + tlsConf := getTLSConfig() + tlsConf.ClientAuth = tls.RequireAndVerifyClientCert + + server, err := quic.Listen(newUDPConnLocalhost(t), tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + + // Usually, the error will occur after the client already finished the handshake. + // However, there's a race condition here. The server's CONNECTION_CLOSE might be + // received before the connection is returned, so we might already get the error while dialing. + if err == nil { + errChan := make(chan error, 1) + go func() { + _, err := conn.AcceptStream(context.Background()) + errChan <- err + }() + + err = <-errChan + } + + require.Error(t, err) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.True(t, transportErr.ErrorCode.IsCryptoError()) + require.Condition(t, func() bool { + errStr := transportErr.Error() + return strings.Contains(errStr, "tls: certificate required") || + strings.Contains(errStr, "tls: bad certificate") + }) +} + +func TestClosedConnectionsInAcceptQueue(t *testing.T) { + dialer := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer dialer.Close() + + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + // Create first connection + conn1, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + conn2, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn2.CloseWithError(0, "") + // close the first connection + const appErrCode quic.ApplicationErrorCode = 12345 + require.NoError(t, conn1.CloseWithError(appErrCode, "")) + + time.Sleep(scaleDuration(25 * time.Millisecond)) // wait for connections to be queued and closed + + // accept all connections, and find the closed one + var closedConn *quic.Conn + for range 2 { + conn, err := server.Accept(ctx) + require.NoError(t, err) + if conn.Context().Err() != nil { + require.Nil(t, closedConn, "only expected a single closed connection") + closedConn = conn + } + } + require.NotNil(t, closedConn, "expected one closed connection") + + _, err = closedConn.AcceptStream(context.Background()) + var appErr *quic.ApplicationError + require.ErrorAs(t, err, &appErr) + require.Equal(t, appErrCode, appErr.ErrorCode) +} + +func TestServerAcceptQueueOverflow(t *testing.T) { + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + dialer := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer dialer.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + // fill up the accept queue + for range protocol.MaxAcceptQueueSize { + conn, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + } + time.Sleep(scaleDuration(25 * time.Millisecond)) // wait for connections to be queued + + // next connection should be rejected + conn, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + _, err = conn.AcceptStream(ctx) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, quic.ConnectionRefused, transportErr.ErrorCode) + + // accept one connection to free up a spot + _, err = server.Accept(ctx) + require.NoError(t, err) + + // should be able to dial again + conn2, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn2.CloseWithError(0, "") + time.Sleep(scaleDuration(25 * time.Millisecond)) + + // but next connection should be rejected again + conn3, err := dialer.Dial(ctx, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + _, err = conn3.AcceptStream(ctx) + require.ErrorAs(t, err, &transportErr) + require.Equal(t, quic.ConnectionRefused, transportErr.ErrorCode) +} + +func TestHandshakeCloseListener(t *testing.T) { + t.Run("using Transport.Listen", func(t *testing.T) { + testHandshakeCloseListener(t, func(tlsConf *tls.Config) *quic.Listener { + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr) + t.Cleanup(func() { tr.Close() }) + + ln, err := tr.Listen(tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + return ln + }) + }) + + t.Run("using Listen", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + testHandshakeCloseListener(t, func(tlsConf *tls.Config) *quic.Listener { + ln, err := quic.Listen(conn, tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + return ln + }) + + // make sure that the Transport didn't close the underlying connection + conn2 := newUDPConnLocalhost(t) + _, err := conn2.WriteTo([]byte("test"), conn2.LocalAddr()) + require.NoError(t, err) + + conn2.SetReadDeadline(time.Now().Add(time.Second)) + b := make([]byte, 1000) + n, err := conn2.Read(b) + require.NoError(t, err) + require.Equal(t, "test", string(b[:n])) + }) + + // This test is somewhat slow (600ms), since the connection entries are kept for 3 PTOs. + t.Run("using ListenAddr", func(t *testing.T) { + var lnAddr *net.UDPAddr + testHandshakeCloseListener(t, func(tlsConf *tls.Config) *quic.Listener { + ln, err := quic.ListenAddr("127.0.0.1:0", tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + lnAddr = ln.Addr().(*net.UDPAddr) + return ln + }) + + // make sure that the Transport closed the underlying connection + if runtime.GOOS != "windows" { // this check doesn't work on Windows + require.Eventually(t, func() bool { + conn, err := net.DialUDP("udp", nil, lnAddr) + require.NoError(t, err) + defer conn.Close() + _, err = conn.Write([]byte("test")) + require.NoError(t, err) + conn.SetReadDeadline(time.Now().Add(scaleDuration(10 * time.Millisecond))) + _, err = conn.Read(make([]byte, 1000)) + require.Error(t, err) + return strings.Contains(err.Error(), "read: connection refused") + }, time.Second, 50*time.Millisecond) + } + }) +} + +func testHandshakeCloseListener(t *testing.T, createListener func(*tls.Config) *quic.Listener) { + connQueued := make(chan struct{}) + var sawFirst atomic.Bool + tlsConf := &tls.Config{ + GetConfigForClient: func(info *tls.ClientHelloInfo) (*tls.Config, error) { + isFirst := sawFirst.CompareAndSwap(false, true) + if isFirst { + } else { + // Sleep for a bit. + // This allows the server to close the connection before the handshake completes. + close(connQueued) + time.Sleep(scaleDuration(10 * time.Millisecond)) + } + return getTLSConfig(), nil + }, + } + + ln := createListener(tlsConf) + + // dial the first connection + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + _, err = ln.Accept(ctx) + require.NoError(t, err) + + errChan := make(chan error, 1) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + errChan <- err + }() + + select { + case <-connQueued: + case <-time.After(scaleDuration(10 * time.Millisecond)): + t.Fatal("timeout waiting for connection queued") + } + + require.NoError(t, ln.Close()) + + select { + case err := <-errChan: + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, quic.ConnectionRefused, transportErr.ErrorCode) + case <-time.After(time.Second): + t.Fatal("timeout waiting for handshaking connection to be rejected") + } + + // the first connection should not be closed + select { + case <-conn.Context().Done(): + t.Fatal("connection was closed") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } +} + +func TestALPN(t *testing.T) { + ln, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + acceptChan := make(chan *quic.Conn, 2) + go func() { + for { + conn, err := ln.Accept(context.Background()) + if err != nil { + return + } + acceptChan <- conn + } + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), nil) + require.NoError(t, err) + cs := conn.ConnectionState() + require.Equal(t, alpn, cs.TLS.NegotiatedProtocol) + + select { + case c := <-acceptChan: + require.Equal(t, alpn, c.ConnectionState().TLS.NegotiatedProtocol) + case <-time.After(time.Second): + t.Fatal("timeout waiting for server connection") + } + require.NoError(t, conn.CloseWithError(0, "")) + + // now try with a different ALPN + tlsConf := getTLSClientConfig() + tlsConf.NextProtos = []string{"foobar"} + ctx, cancel = context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err = quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), tlsConf, nil) + require.Error(t, err) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.True(t, transportErr.ErrorCode.IsCryptoError()) + require.Contains(t, transportErr.Error(), "no application protocol") +} + +func TestTokensFromNewTokenFrames(t *testing.T) { + t.Run("MaxTokenAge: 1 hour", func(t *testing.T) { + testTokensFromNewTokenFrames(t, 0, true) + }) + // If unset, the default value is 24h. + t.Run("MaxTokenAge: default", func(t *testing.T) { + testTokensFromNewTokenFrames(t, 0, true) + }) + t.Run("MaxTokenAge: very short", func(t *testing.T) { + testTokensFromNewTokenFrames(t, time.Microsecond, false) + }) +} + +func testTokensFromNewTokenFrames(t *testing.T, maxTokenAge time.Duration, expectTokenUsed bool) { + addrVerifiedChan := make(chan bool, 2) + quicConf := getQuicConfig(nil) + quicConf.GetConfigForClient = func(info *quic.ClientInfo) (*quic.Config, error) { + addrVerifiedChan <- info.AddrVerified + return quicConf, nil + } + tr := &quic.Transport{Conn: newUDPConnLocalhost(t), MaxTokenAge: maxTokenAge} + addTracer(tr) + defer tr.Close() + server, err := tr.Listen(getTLSConfig(), quicConf) + require.NoError(t, err) + defer server.Close() + + // dial the first connection and receive the token + acceptChan := make(chan error, 2) + go func() { + _, err := server.Accept(context.Background()) + acceptChan <- err + _, err = server.Accept(context.Background()) + acceptChan <- err + }() + + gets := make(chan string, 2) + puts := make(chan string, 2) + ts := newTokenStore(gets, puts) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(&quic.Config{TokenStore: ts})) + require.NoError(t, err) + + // verify token store was used + select { + case <-gets: + case <-time.After(time.Second): + t.Fatal("timeout waiting for token store get") + } + select { + case <-puts: + case <-time.After(time.Second): + t.Fatal("timeout waiting for token store put") + } + select { + case addrVerified := <-addrVerifiedChan: + require.False(t, addrVerified) + case <-time.After(time.Second): + t.Fatal("timeout waiting for addr verified") + } + select { + case <-acceptChan: + case <-time.After(time.Second): + t.Fatal("timeout waiting for accept") + } + // received a token. Close this connection. + require.NoError(t, conn.CloseWithError(0, "")) + + time.Sleep(scaleDuration(5 * time.Millisecond)) + conn, err = quic.Dial(ctx, newUDPConnLocalhost(t), server.Addr(), getTLSClientConfig(), getQuicConfig(&quic.Config{TokenStore: ts})) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + select { + case addrVerified := <-addrVerifiedChan: + // this time, the address was verified using the token + if expectTokenUsed { + require.True(t, addrVerified) + } else { + require.False(t, addrVerified) + } + + case <-time.After(time.Second): + t.Fatal("timeout waiting for addr verified") + } + select { + case <-gets: + case <-time.After(time.Second): + t.Fatal("timeout waiting for token store get") + } + select { + case <-acceptChan: + case <-time.After(time.Second): + t.Fatal("timeout waiting for accept") + } +} + +func TestInvalidToken(t *testing.T) { + const rtt = 10 * time.Millisecond + + // The validity period of the retry token is the handshake timeout, + // which is twice the handshake idle timeout. + // By setting the handshake timeout shorter than the RTT, the token will have + // expired by the time it reaches the server. + serverConfig := getQuicConfig(&quic.Config{HandshakeIdleTimeout: rtt / 5}) + + tr := &quic.Transport{ + Conn: newUDPConnLocalhost(t), + VerifySourceAddress: func(net.Addr) bool { return true }, + } + addTracer(tr) + defer tr.Close() + + server, err := tr.Listen(getTLSConfig(), serverConfig) + require.NoError(t, err) + defer server.Close() + + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: server.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { return rtt / 2 }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err = quic.Dial(ctx, newUDPConnLocalhost(t), proxy.LocalAddr(), getTLSClientConfig(), nil) + require.Error(t, err) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, quic.InvalidToken, transportErr.ErrorCode) +} + +func TestGetConfigForClient(t *testing.T) { + var calledFrom net.Addr + serverConfig := getQuicConfig(&quic.Config{EnableDatagrams: true}) + serverConfig.GetConfigForClient = func(info *quic.ClientInfo) (*quic.Config, error) { + conf := serverConfig.Clone() + conf.EnableDatagrams = true + calledFrom = info.RemoteAddr + return getQuicConfig(conf), nil + } + ln, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), serverConfig) + require.NoError(t, err) + + acceptDone := make(chan struct{}) + go func() { + _, err := ln.Accept(context.Background()) + require.NoError(t, err) + close(acceptDone) + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(&quic.Config{EnableDatagrams: true})) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + cs := conn.ConnectionState() + require.True(t, cs.SupportsDatagrams.Remote, "server should advertise datagram support") + require.True(t, cs.SupportsDatagrams.Local, "client should have datagram support enabled") + + select { + case <-acceptDone: + case <-time.After(time.Second): + t.Fatal("timeout waiting for accept") + } + + require.NoError(t, ln.Close()) + require.Equal(t, conn.LocalAddr().(*net.UDPAddr).Port, calledFrom.(*net.UDPAddr).Port) +} + +func TestGetConfigForClientErrorsConnectionRejection(t *testing.T) { + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + GetConfigForClient: func(info *quic.ClientInfo) (*quic.Config, error) { + return nil, errors.New("rejected") + }, + }), + ) + require.NoError(t, err) + + acceptChan := make(chan bool, 1) + go func() { + _, err := ln.Accept(context.Background()) + acceptChan <- err == nil + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err = quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.ConnectionRefused, transportErr.ErrorCode) + + // verify no connection was accepted + ln.Close() + require.False(t, <-acceptChan) +} + +func TestNoPacketsSentWhenClientHelloFails(t *testing.T) { + conn := newUDPConnLocalhost(t) + + packetChan := make(chan struct{}, 1) + go func() { + for { + _, _, err := conn.ReadFromUDP(make([]byte, protocol.MaxPacketBufferSize)) + if err != nil { + return + } + select { + case packetChan <- struct{}{}: + default: + } + } + }() + + tlsConf := getTLSClientConfig() + tlsConf.NextProtos = []string{""} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err := quic.Dial(ctx, newUDPConnLocalhost(t), conn.LocalAddr(), tlsConf, getQuicConfig(nil)) + + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.True(t, transportErr.ErrorCode.IsCryptoError()) + require.Contains(t, err.Error(), "tls: invalid NextProtos value") + + // verify no packets were sent + select { + case <-packetChan: + t.Fatal("received unexpected packet") + case <-time.After(50 * time.Millisecond): + // no packets received, as expected + } +} + +func TestServerTransportClose(t *testing.T) { + tlsServerConf := getTLSConfig() + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + server, err := tr.Listen(tlsServerConf, getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + // the first conn is accepted by the server... + conn1, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIdleTimeout: scaleDuration(50 * time.Millisecond)}), + ) + require.NoError(t, err) + + sconn, err := server.Accept(ctx) + require.NoError(t, err) + require.Equal(t, conn1.LocalAddr(), sconn.RemoteAddr()) + + // ...the second conn isn't, it remains in the server's accept queue + conn2, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIdleTimeout: scaleDuration(50 * time.Millisecond)}), + ) + require.NoError(t, err) + + time.Sleep(scaleDuration(10 * time.Millisecond)) + + // closing the Transport abruptly terminates connections + require.NoError(t, tr.Close()) + + select { + case <-sconn.Context().Done(): + require.ErrorIs(t, context.Cause(sconn.Context()), quic.ErrTransportClosed) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // no CONNECTION_CLOSE frame is sent to the peers + select { + case <-conn1.Context().Done(): + require.ErrorIs(t, context.Cause(conn1.Context()), &quic.IdleTimeoutError{}) + case <-time.After(time.Second): + t.Fatal("timeout") + } + select { + case <-conn2.Context().Done(): + require.ErrorIs(t, context.Cause(conn1.Context()), &quic.IdleTimeoutError{}) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // Accept should error after the transport was closed + ctx, cancel = context.WithTimeout(context.Background(), time.Second) + defer cancel() + accepted, err := server.Accept(ctx) + require.ErrorIs(t, err, quic.ErrTransportClosed) + require.Nil(t, accepted) +} diff --git a/third_party/quic-go/integrationtests/self/http_datagram_test.go b/third_party/quic-go/integrationtests/self/http_datagram_test.go new file mode 100644 index 0000000..8a318ba --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_datagram_test.go @@ -0,0 +1,323 @@ +package self_test + +import ( + "bytes" + "context" + "encoding/binary" + "fmt" + "io" + "net" + "net/http" + "net/url" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + + "github.com/stretchr/testify/require" +) + +func TestHTTPSettings(t *testing.T) { + mux := http.NewServeMux() + port := startHTTPServer(t, mux) + + t.Run("server settings", func(t *testing.T) { + tlsConf := getTLSClientConfig() + tlsConf.NextProtos = []string{http3.NextProtoH3} + conn, err := quic.Dial( + context.Background(), + newUDPConnLocalhost(t), + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: port}, + tlsConf, + getQuicConfig(nil), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + var tr http3.Transport + cc := tr.NewClientConn(conn) + + select { + case <-cc.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("didn't receive HTTP/3 settings") + } + + settings := cc.Settings() + require.True(t, settings.EnableExtendedConnect) + require.False(t, settings.EnableDatagrams) + require.Empty(t, settings.Other) + }) + + t.Run("client settings", func(t *testing.T) { + connChan := make(chan http3.Settingser, 1) + mux.HandleFunc("/settings", func(w http.ResponseWriter, r *http.Request) { + connChan <- w.(http3.Settingser) + w.WriteHeader(http.StatusOK) + }) + + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + QUICConfig: getQuicConfig(&quic.Config{ + MaxIdleTimeout: 10 * time.Second, + EnableDatagrams: true, + }), + EnableDatagrams: true, + AdditionalSettings: map[uint64]uint64{1337: 42}, + } + addDialCallback(t, tr) + defer tr.Close() + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://localhost:%d/settings", port), nil) + require.NoError(t, err) + + _, err = tr.RoundTrip(req) + require.NoError(t, err) + var conn http3.Settingser + select { + case conn = <-connChan: + case <-time.After(time.Second): + t.Fatal("didn't receive HTTP/3 connection") + } + + select { + case <-conn.ReceivedSettings(): + case <-time.After(time.Second): + t.Fatal("didn't receive HTTP/3 settings") + } + settings := conn.Settings() + require.NotNil(t, settings) + require.True(t, settings.EnableDatagrams) + require.False(t, settings.EnableExtendedConnect) + require.Equal(t, uint64(42), settings.Other[1337]) + }) +} + +func dialAndOpenHTTPDatagramStream(t *testing.T, addr string) *http3.RequestStream { + t.Helper() + + u, err := url.Parse(addr) + require.NoError(t, err) + + tlsConf := getTLSClientConfig() + tlsConf.NextProtos = []string{http3.NextProtoH3} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverAddr, err := net.ResolveUDPAddr("udp4", u.Host) + require.NoError(t, err) + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + serverAddr, + tlsConf, + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + t.Cleanup(func() { conn.CloseWithError(0, "") }) + + tr := http3.Transport{EnableDatagrams: true} + t.Cleanup(func() { tr.Close() }) + cc := tr.NewClientConn(conn) + t.Cleanup(func() { cc.CloseWithError(0, "") }) + str, err := cc.OpenRequestStream(ctx) + require.NoError(t, err) + req := &http.Request{ + Method: http.MethodConnect, + Proto: "datagrams", + Host: u.Host, + URL: u, + } + require.NoError(t, str.SendRequestHeader(req)) + + rsp, err := str.ReadResponse() + require.NoError(t, err) + require.Equal(t, http.StatusOK, rsp.StatusCode) + return str +} + +func TestHTTPDatagrams(t *testing.T) { + errChan := make(chan error, 1) + const num = 5 + datagramChan := make(chan struct{}, num) + mux := http.NewServeMux() + mux.HandleFunc("/datagrams", func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodConnect { + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + s := w.(http3.Settingser) + select { + case <-s.ReceivedSettings(): + case <-time.After(time.Second): + w.WriteHeader(http.StatusBadRequest) + return + } + if !s.Settings().EnableDatagrams { + w.WriteHeader(http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusOK) + + str := w.(http3.HTTPStreamer).HTTPStream() + go str.Read([]byte{0}) // need to continue reading from stream to observe state transitions + + for { + if _, err := str.ReceiveDatagram(context.Background()); err != nil { + errChan <- err + return + } + datagramChan <- struct{}{} + } + }) + + port := startHTTPServer(t, mux, func(s *http3.Server) { s.EnableDatagrams = true }) + str := dialAndOpenHTTPDatagramStream(t, fmt.Sprintf("https://localhost:%d/datagrams", port)) + + for i := range num { + b := make([]byte, 8) + binary.BigEndian.PutUint64(b, uint64(i)) + require.NoError(t, str.SendDatagram(bytes.Repeat(b, 100))) + } + var count int +loop: + for { + select { + case <-datagramChan: + count++ + if count >= num*4/5 { + break loop + } + case err := <-errChan: + t.Fatalf("receiving datagrams failed: %s", err) + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + str.CancelWrite(42) + + select { + case err := <-errChan: + var serr *quic.StreamError + require.ErrorAs(t, err, &serr) + require.Equal(t, quic.StreamErrorCode(42), serr.ErrorCode) + case <-time.After(time.Second): + t.Fatal("didn't receive error") + } +} + +func TestHTTPDatagramClose(t *testing.T) { + errChan := make(chan error, 1) + datagramChan := make(chan []byte, 1) + mux := http.NewServeMux() + mux.HandleFunc("/datagrams", func(w http.ResponseWriter, r *http.Request) { + s := w.(http3.Settingser) + select { + case <-s.ReceivedSettings(): + case <-time.After(time.Second): + w.WriteHeader(http.StatusBadRequest) + return + } + if !s.Settings().EnableDatagrams { + w.WriteHeader(http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusOK) + + str := w.(http3.HTTPStreamer).HTTPStream() + go str.Read([]byte{0}) // need to continue reading from stream to observe state transitions + + for { + data, err := str.ReceiveDatagram(context.Background()) + if err != nil { + errChan <- err + return + } + datagramChan <- data + } + }) + + port := startHTTPServer(t, mux, func(s *http3.Server) { s.EnableDatagrams = true }) + str := dialAndOpenHTTPDatagramStream(t, fmt.Sprintf("https://localhost:%d/datagrams", port)) + go str.Read([]byte{0}) + + require.NoError(t, str.SendDatagram([]byte("foo"))) + select { + case data := <-datagramChan: + require.Equal(t, []byte("foo"), data) + case <-time.After(time.Second): + t.Fatal("didn't receive datagram") + } + // signal that we're done sending + str.Close() + + var resetErr error + select { + case resetErr = <-errChan: + case <-time.After(time.Second): + t.Fatal("didn't receive error") + } + require.Equal(t, io.EOF, resetErr) + + // make sure we can't send anymore + require.Error(t, str.SendDatagram([]byte("foo"))) +} + +func TestHTTPDatagramStreamReset(t *testing.T) { + errChan := make(chan error, 1) + datagramChan := make(chan []byte, 1) + mux := http.NewServeMux() + mux.HandleFunc("/datagrams", func(w http.ResponseWriter, r *http.Request) { + s := w.(http3.Settingser) + select { + case <-s.ReceivedSettings(): + case <-time.After(time.Second): + w.WriteHeader(http.StatusBadRequest) + return + } + if !s.Settings().EnableDatagrams { + w.WriteHeader(http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusOK) + + str := w.(http3.HTTPStreamer).HTTPStream() + go str.Read([]byte{0}) // need to continue reading from stream to observe state transitions + + for { + data, err := str.ReceiveDatagram(context.Background()) + if err != nil { + errChan <- err + return + } + str.CancelRead(42) + datagramChan <- data + } + }) + + port := startHTTPServer(t, mux, func(s *http3.Server) { s.EnableDatagrams = true }) + str := dialAndOpenHTTPDatagramStream(t, fmt.Sprintf("https://localhost:%d/datagrams", port)) + go str.Read([]byte{0}) + + require.NoError(t, str.SendDatagram([]byte("foo"))) + select { + case data := <-datagramChan: + require.Equal(t, []byte("foo"), data) + case <-time.After(time.Second): + t.Fatal("didn't receive datagram") + } + + var resetErr error + select { + case resetErr = <-errChan: + case <-time.After(time.Second): + t.Fatal("didn't receive error") + } + require.Equal(t, &quic.StreamError{ErrorCode: 42, Remote: false}, resetErr) + + var err error + require.Eventually(t, func() bool { + err = str.SendDatagram([]byte("foo")) + return err != nil + }, time.Second, 10*time.Millisecond) + // make sure we can't send anymore + require.Equal(t, &quic.StreamError{ErrorCode: 42, Remote: true}, err) +} diff --git a/third_party/quic-go/integrationtests/self/http_hotswap_test.go b/third_party/quic-go/integrationtests/self/http_hotswap_test.go new file mode 100644 index 0000000..1a3d2de --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_hotswap_test.go @@ -0,0 +1,111 @@ +package self_test + +import ( + "io" + "net" + "net/http" + "strconv" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/stretchr/testify/require" +) + +func TestHTTP3ServerHotswap(t *testing.T) { + mux1 := http.NewServeMux() + mux1.HandleFunc("/hello1", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World 1!\n") // don't check the error here. Stream may be reset. + }) + + mux2 := http.NewServeMux() + mux2.HandleFunc("/hello2", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World 2!\n") // don't check the error here. Stream may be reset. + }) + + server1 := &http3.Server{ + Handler: mux1, + QUICConfig: getQuicConfig(nil), + } + server2 := &http3.Server{ + Handler: mux2, + QUICConfig: getQuicConfig(nil), + } + + tlsConf := http3.ConfigureTLSConfig(getTLSConfig()) + ln, err := quic.ListenEarly(newUDPConnLocalhost(t), tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + port := strconv.Itoa(ln.Addr().(*net.UDPAddr).Port) + + newClient := func() *http.Client { + return &http.Client{ + Transport: &http3.Transport{ + TLSClientConfig: getTLSClientConfig(), + DisableCompression: true, + QUICConfig: getQuicConfig(&quic.Config{MaxIdleTimeout: 10 * time.Second}), + }, + } + } + + client := newClient() + + defer func() { + require.NoError(t, ln.Close()) + }() + + // open first server and make single request to it + errChan1 := make(chan error, 1) + go func() { errChan1 <- server1.ServeListener(ln) }() + + resp, err := client.Get("https://localhost:" + port + "/hello1") + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, "Hello, World 1!\n", string(body)) + + // open second server with same underlying listener + errChan2 := make(chan error, 1) + go func() { errChan2 <- server2.ServeListener(ln) }() + + time.Sleep(scaleDuration(20 * time.Millisecond)) + select { + case err := <-errChan1: + t.Fatalf("server1 stopped unexpectedly: %v", err) + case err := <-errChan2: + t.Fatalf("server2 stopped unexpectedly: %v", err) + default: + } + + // now close first server + require.NoError(t, server1.Close()) + select { + case err := <-errChan1: + require.ErrorIs(t, err, http.ErrServerClosed) + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for server1 to stop") + } + require.NoError(t, client.Transport.(*http3.Transport).Close()) + client = newClient() + defer func() { + require.NoError(t, client.Transport.(*http3.Transport).Close()) + }() + + // verify that new connections are handled by the second server now + resp, err = client.Get("https://localhost:" + port + "/hello2") + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err = io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, "Hello, World 2!\n", string(body)) + + // close the other server + require.NoError(t, server2.Close()) + select { + case err := <-errChan2: + require.ErrorIs(t, err, http.ErrServerClosed) + case <-time.After(time.Second): + t.Fatal("timed out waiting for server2 to stop") + } +} diff --git a/third_party/quic-go/integrationtests/self/http_qlog_test.go b/third_party/quic-go/integrationtests/self/http_qlog_test.go new file mode 100644 index 0000000..78770c5 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_qlog_test.go @@ -0,0 +1,83 @@ +package self_test + +import ( + "context" + "fmt" + "io" + "net" + "net/http" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + h3qlog "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestHTTP3Qlog(t *testing.T) { + serverTrace := newMockTrace() + clientTrace := newMockTrace() + + mux := http.NewServeMux() + mux.HandleFunc("/hello", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World!\n") + }) + + server := &http3.Server{ + Handler: mux, + TLSConfig: getTLSConfig(), + QUICConfig: getQuicConfig(&quic.Config{ + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return serverTrace + }, + }), + } + + conn := newUDPConnLocalhost(t) + done := make(chan struct{}) + go func() { + defer close(done) + server.Serve(conn) + }() + port := conn.LocalAddr().(*net.UDPAddr).Port + + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + QUICConfig: getQuicConfig(&quic.Config{ + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return clientTrace + }, + }), + } + addDialCallback(t, tr) + cl := &http.Client{Transport: tr} + + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, "Hello, World!\n", string(body)) + resp.Body.Close() + + assert.Equal(t, 2, clientTrace.OpenRecorders()) + assert.Equal(t, 2, serverTrace.OpenRecorders()) + + tr.Close() + server.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("server didn't shut down") + } + + // Recorders are closed in an AfterFunc, so we need to wait for them to be closed. + assert.Eventually(t, func() bool { return clientTrace.OpenRecorders() == 0 }, time.Second, 10*time.Millisecond, "client recorders should be closed") + assert.Eventually(t, func() bool { return serverTrace.OpenRecorders() == 0 }, time.Second, 10*time.Millisecond, "server recorders should be closed") + assert.Equal(t, []string{h3qlog.EventSchema}, clientTrace.SchemasChecked) + assert.Equal(t, []string{h3qlog.EventSchema}, serverTrace.SchemasChecked) +} diff --git a/third_party/quic-go/integrationtests/self/http_raw_conn_test.go b/third_party/quic-go/integrationtests/self/http_raw_conn_test.go new file mode 100644 index 0000000..b77b089 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_raw_conn_test.go @@ -0,0 +1,171 @@ +package self_test + +import ( + "context" + "fmt" + "io" + "net" + "net/http" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +// This test tests the HTTP/3 raw connection functionality, +// which is primarily used by WebTransport. +func TestHTTPRawConn(t *testing.T) { + const magicValue = 0x123456 + + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + + clientPacketConn, serverPacketConn, closeFn := newSimnetLink(t, rtt) + defer closeFn(t) + + ln, err := quic.ListenEarly( + serverPacketConn, + http3.ConfigureTLSConfig(getTLSConfig()), + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + defer ln.Close() + + start := time.Now() + + mux := http.NewServeMux() + mux.HandleFunc("/data", func(w http.ResponseWriter, r *http.Request) { w.Write(PRData) }) + server := &http3.Server{ + Handler: mux, + EnableDatagrams: true, + } + defer server.Close() + + // run the server in a separate Goroutine, so we can make sure that SETTINGS are sent in 0.5-RTT data + errChan := make(chan error, 1) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverConn, err := ln.Accept(ctx) + if err != nil { + errChan <- err + return + } + + rawServerConn, err := server.NewRawServerConn(serverConn) + if err != nil { + errChan <- err + return + } + var wg sync.WaitGroup + // accept and handle unidirectional streams opened by the client + wg.Go(func() { + for { + str, err := serverConn.AcceptUniStream(context.Background()) + if err != nil { + return + } + go rawServerConn.HandleUnidirectionalStream(str) + } + }) + // accept and handle bidirectional streams opened by the client + wg.Go(func() { + for { + str, err := serverConn.AcceptStream(context.Background()) + if err != nil { + return + } + v, _ := quicvarint.Peek(str) + if v == magicValue { + go func() { + // read the previously peeked value + quicvarint.Read(quicvarint.NewReader(str)) + defer str.Close() + io.Copy(str, str) + }() + } else { + go rawServerConn.HandleRequestStream(str) + } + } + }) + wg.Wait() + <-serverConn.Context().Done() + errChan <- nil + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Hour) + defer cancel() + clientConn, err := quic.Dial( + ctx, + clientPacketConn, + serverPacketConn.LocalAddr(), + http3.ConfigureTLSConfig(getTLSClientConfig()), + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + defer clientConn.CloseWithError(0, "") + + tr := &http3.Transport{ + EnableDatagrams: true, + } + rawClientConn := tr.NewRawClientConn(clientConn) + // accept and handle unidirectional streams opened by the server + go func() { + for { + str, err := clientConn.AcceptUniStream(ctx) + if err != nil { + return + } + go rawClientConn.HandleUnidirectionalStream(str) + } + }() + + select { + case <-rawClientConn.ReceivedSettings(): + settings := rawClientConn.Settings() + require.True(t, settings.EnableDatagrams) + // the server sends SETTINGS in 0.5-RTT data, so they should be received after 1 RTT + require.Equal(t, rtt, time.Since(start)) + case <-time.After(time.Second): + t.Fatal("timeout waiting for HTTP/3 settings") + } + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://%s/data", serverPacketConn.LocalAddr().(*net.UDPAddr)), nil) + require.NoError(t, err) + reqStr, err := rawClientConn.OpenRequestStream(ctx) + require.NoError(t, err) + require.NoError(t, reqStr.SendRequestHeader(req)) + resp, err := reqStr.ReadResponse() + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + data, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, PRData, data) + require.NoError(t, resp.Body.Close()) + + str, err := clientConn.OpenStream() + require.NoError(t, err) + b := quicvarint.Append(nil, magicValue) + b = append(b, []byte("lorem ipsum dolor sit amet")...) + _, err = str.Write(b) + require.NoError(t, err) + require.NoError(t, str.Close()) + data, err = io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, []byte("lorem ipsum dolor sit amet"), data) + + clientConn.CloseWithError(0, "") + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout waiting for server to close") + } + }) +} diff --git a/third_party/quic-go/integrationtests/self/http_shutdown_test.go b/third_party/quic-go/integrationtests/self/http_shutdown_test.go new file mode 100644 index 0000000..3a0d280 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_shutdown_test.go @@ -0,0 +1,520 @@ +package self_test + +import ( + "context" + "crypto/tls" + "fmt" + "io" + "net" + "net/http" + "net/url" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestHTTPShutdown(t *testing.T) { + mux := http.NewServeMux() + var server *http3.Server + port := startHTTPServer(t, mux, func(s *http3.Server) { server = s }) + client := newHTTP3Client(t) + + mux.HandleFunc("/shutdown", func(w http.ResponseWriter, r *http.Request) { + go func() { + require.NoError(t, server.Close()) + }() + time.Sleep(scaleDuration(10 * time.Millisecond)) // make sure the server started shutting down + }) + + _, err := client.Get(fmt.Sprintf("https://localhost:%d/shutdown", port)) + require.Error(t, err) + var appErr *http3.Error + require.ErrorAs(t, err, &appErr) + require.Equal(t, http3.ErrCodeNoError, appErr.ErrorCode) +} + +func TestGracefulShutdownShortRequest(t *testing.T) { + var server *http3.Server + mux := http.NewServeMux() + port := startHTTPServer(t, mux, func(s *http3.Server) { server = s }) + errChan := make(chan error, 1) + proceed := make(chan struct{}) + mux.HandleFunc("/shutdown", func(w http.ResponseWriter, r *http.Request) { + go func() { + defer close(errChan) + errChan <- server.Shutdown(context.Background()) + }() + w.WriteHeader(http.StatusOK) + w.(http.Flusher).Flush() + <-proceed + w.Write([]byte("shutdown")) + }) + + connChan := make(chan *quic.Conn, 1) + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + Dial: func(ctx context.Context, a string, tlsConf *tls.Config, conf *quic.Config) (*quic.Conn, error) { + addr, err := net.ResolveUDPAddr("udp", a) + if err != nil { + return nil, err + } + conn, err := quic.DialEarly(ctx, newUDPConnLocalhost(t), addr, tlsConf, conf) + connChan <- conn + return conn, err + }, + } + t.Cleanup(func() { tr.Close() }) + + client := &http.Client{Transport: tr} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/shutdown", port), nil) + require.NoError(t, err) + resp, err := client.Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + var conn *quic.Conn + select { + case conn = <-connChan: + default: + t.Fatal("expected a connection") + } + + type result struct { + body []byte + err error + } + resultChan := make(chan result, 1) + go func() { + body, err := io.ReadAll(resp.Body) + resultChan <- result{body: body, err: err} + }() + select { + case <-resultChan: + t.Fatal("request body shouldn't have been read yet") + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + select { + case <-conn.Context().Done(): + t.Fatal("connection shouldn't have been closed") + default: + } + + // allow the request to proceed + close(proceed) + select { + case res := <-resultChan: + require.NoError(t, res.err) + require.Equal(t, []byte("shutdown"), res.body) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // now that the stream count dropped to 0, the client should close the connection + select { + case <-conn.Context().Done(): + var appErr *quic.ApplicationError + require.ErrorAs(t, context.Cause(conn.Context()), &appErr) + assert.False(t, appErr.Remote) + assert.Equal(t, quic.ApplicationErrorCode(http3.ErrCodeNoError), appErr.ErrorCode) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("shutdown did not complete") + } +} + +func TestGracefulShutdownIdleConnection(t *testing.T) { + var server *http3.Server + port := startHTTPServer(t, http.NewServeMux(), func(s *http3.Server) { server = s }) + + connChan := make(chan *quic.Conn, 1) + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + Dial: func(ctx context.Context, a string, tlsConf *tls.Config, conf *quic.Config) (*quic.Conn, error) { + addr, err := net.ResolveUDPAddr("udp", a) + if err != nil { + return nil, err + } + conn, err := quic.DialEarly(ctx, newUDPConnLocalhost(t), addr, tlsConf, conf) + connChan <- conn + return conn, err + }, + } + t.Cleanup(func() { tr.Close() }) + + client := &http.Client{Transport: tr} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/", port), nil) + require.NoError(t, err) + resp, err := client.Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusNotFound, resp.StatusCode) + require.NoError(t, resp.Body.Close()) + + var conn *quic.Conn + select { + case conn = <-connChan: + default: + t.Fatal("expected a connection") + } + // the connection should still be alive (and idle) + select { + case <-conn.Context().Done(): + t.Fatal("connection shouldn't have been closed") + default: + } + + shutdownChan := make(chan error, 1) + go func() { shutdownChan <- server.Shutdown(context.Background()) }() + + // since the connection is idle, the client should close it immediately + select { + case <-conn.Context().Done(): + var appErr *quic.ApplicationError + require.ErrorAs(t, context.Cause(conn.Context()), &appErr) + assert.False(t, appErr.Remote) + assert.Equal(t, quic.ApplicationErrorCode(http3.ErrCodeNoError), appErr.ErrorCode) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestGracefulShutdownLongLivedRequest(t *testing.T) { + delay := scaleDuration(25 * time.Millisecond) + errChan := make(chan error, 1) + requestChan := make(chan time.Duration, 1) + + var server *http3.Server + mux := http.NewServeMux() + port := startHTTPServer(t, mux, func(s *http3.Server) { server = s }) + mux.HandleFunc("/shutdown", func(w http.ResponseWriter, r *http.Request) { + start := time.Now() + w.WriteHeader(http.StatusOK) + w.(http.Flusher).Flush() + + // The request simulated here takes longer than the server's graceful shutdown period. + // We expect it to be terminated once the server shuts down. + go func() { + ctx, cancel := context.WithTimeout(context.Background(), delay) + defer cancel() + errChan <- server.Shutdown(ctx) + }() + + // measure how long it takes until the request errors + for t := range time.NewTicker(delay / 10).C { + if _, err := w.Write([]byte(t.String())); err != nil { + requestChan <- time.Since(start) + return + } + } + }) + + start := time.Now() + resp, err := newHTTP3Client(t).Get(fmt.Sprintf("https://localhost:%d/shutdown", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + _, err = io.Copy(io.Discard, resp.Body) + require.Error(t, err) + var h3Err *http3.Error + require.ErrorAs(t, err, &h3Err) + require.Equal(t, http3.ErrCodeNoError, h3Err.ErrorCode) + took := time.Since(start) + require.InDelta(t, delay.Seconds(), took.Seconds(), (delay / 2).Seconds()) + + // make sure that shutdown returned due to context deadline + select { + case err := <-errChan: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(time.Second): + t.Fatal("shutdown did not return due to context deadline") + } + + select { + case requestDuration := <-requestChan: + require.InDelta(t, delay.Seconds(), requestDuration.Seconds(), (delay / 2).Seconds()) + case <-time.After(time.Second): + t.Fatal("did not receive request duration") + } +} + +func TestGracefulShutdownPendingStreams(t *testing.T) { + rtt := scaleDuration(25 * time.Millisecond) + + handlerChan := make(chan struct{}, 1) + mux := http.NewServeMux() + mux.HandleFunc("/helloworld", func(w http.ResponseWriter, r *http.Request) { + handlerChan <- struct{}{} + time.Sleep(rtt) + w.Write([]byte("hello world")) + }) + var server *http3.Server + port := startHTTPServer(t, mux, func(s *http3.Server) { server = s }) + connChan := make(chan *quic.Conn, 1) + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + Dial: func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + a, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return nil, err + } + conn, err := quic.DialEarly(ctx, newUDPConnLocalhost(t), a, tlsCfg, cfg) + connChan <- conn + return conn, err + }, + } + cl := &http.Client{Transport: tr} + + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: port}, + DelayPacket: func(_ quicproxy.Direction, _, _ net.Addr, _ []byte) time.Duration { return rtt }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + errChan := make(chan error, 1) + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://%s/helloworld", proxy.LocalAddr()), nil) + require.NoError(t, err) + go func() { + resp, err := cl.Do(req) + if err != nil { + errChan <- err + return + } + if resp.StatusCode != http.StatusOK { + errChan <- fmt.Errorf("expected status code %d, got %d", http.StatusOK, resp.StatusCode) + } + }() + + select { + case <-handlerChan: + case <-time.After(time.Second): + t.Fatal("did not receive request") + } + + shutdownChan := make(chan error, 1) + ctx, cancel := context.WithCancel(context.Background()) + go func() { shutdownChan <- server.Shutdown(ctx) }() + time.Sleep(rtt / 2) // wait for the server to start shutting down + + var conn *quic.Conn + select { + case conn = <-connChan: + case <-time.After(time.Second): + t.Fatal("connection was not opened") + } + + // make sure that the server rejects further requests + for range 3 { + str, err := conn.OpenStreamSync(ctx) + require.NoError(t, err) + str.Write([]byte("foobar")) + select { + case <-str.Context().Done(): + case <-time.After(time.Second): + t.Fatal("stream was not rejected") + } + _, err = str.Read(make([]byte, 10)) + var serr *quic.StreamError + require.ErrorAs(t, err, &serr) + require.Equal(t, quic.StreamErrorCode(http3.ErrCodeRequestRejected), serr.ErrorCode) + } + + cancel() + select { + case err := <-shutdownChan: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("shutdown did not complete") + } +} + +func TestHTTP3ListenerClosing(t *testing.T) { + t.Run("application listener", func(t *testing.T) { + testHTTP3ListenerClosing(t, false, true) + }) + t.Run("listener created by the http3.Server", func(t *testing.T) { + testHTTP3ListenerClosing(t, false, false) + }) +} + +func TestHTTP3ListenerGracefulShutdown(t *testing.T) { + t.Run("application listener", func(t *testing.T) { + testHTTP3ListenerClosing(t, true, true) + }) + t.Run("listener created by the http3.Server", func(t *testing.T) { + testHTTP3ListenerClosing(t, true, false) + }) +} + +func testHTTP3ListenerClosing(t *testing.T, graceful, useApplicationListener bool) { + dial := func(t *testing.T, ctx context.Context, u *url.URL) error { + t.Helper() + tlsConf := getTLSClientConfig() + tlsConf.NextProtos = []string{http3.NextProtoH3} + tr := &http3.Transport{TLSClientConfig: tlsConf} + defer tr.Close() + addDialCallback(t, tr) + cl := &http.Client{Transport: tr} + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil) + require.NoError(t, err) + resp, err := cl.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) + return nil + } + + mux := http.NewServeMux() + mux.HandleFunc("/ok", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + handlerChan := make(chan struct{}) + mux.HandleFunc("/long", func(w http.ResponseWriter, r *http.Request) { + <-handlerChan + w.WriteHeader(http.StatusOK) + }) + + tlsConf := http3.ConfigureTLSConfig(getTLSConfig()) + server := &http3.Server{ + Handler: mux, + // the following values will be ignored when using ServeListener + TLSConfig: tlsConf, + QUICConfig: getQuicConfig(nil), + Addr: "127.0.0.1:0", + } + + serveChan := make(chan error, 1) + var host string + var ln *quic.EarlyListener // only set when using application listener + if useApplicationListener { + var err error + ln, err = quic.ListenEarly(newUDPConnLocalhost(t), tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + host = ln.Addr().String() + go func() { serveChan <- server.ServeListener(ln) }() + } else { + go func() { serveChan <- server.ListenAndServe() }() + // The server is listening on a random port, and the only way to get the port + // is to parse the Alt-Svc header. + var port int + require.Eventually(t, func() bool { + hdr := make(http.Header) + server.SetQUICHeaders(hdr) + altSvc := hdr.Get("Alt-Svc") + n, err := fmt.Sscanf(altSvc, `h3=":%d"`, &port) + return err == nil && n == 1 + }, time.Second, 10*time.Millisecond) + host = fmt.Sprintf("127.0.0.1:%d", port) + } + + u := &url.URL{Scheme: "https", Host: host, Path: "/ok"} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + require.NoError(t, dial(t, ctx, u)) + + longReqChan := make(chan error, 1) + shutdownChan := make(chan error, 1) + if graceful { + go func() { + u := &url.URL{Scheme: "https", Host: host, Path: "/long"} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + longReqChan <- dial(t, ctx, u) + }() + time.Sleep(scaleDuration(10 * time.Millisecond)) + + go func() { shutdownChan <- server.Shutdown(context.Background()) }() + } else { + require.NoError(t, server.Close()) + } + + select { + case err := <-serveChan: + require.ErrorIs(t, err, http.ErrServerClosed) + case <-time.After(time.Second): + t.Fatal("server did not stop") + } + + // If the listener was created by the http3.Server, it will now be closed. + if !useApplicationListener { + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(10*time.Millisecond)) + defer cancel() + require.ErrorIs(t, dial(t, ctx, u), context.DeadlineExceeded) + } else { + // If the listener was created by the application, it will not be closed, + // and it can be used to accept new connections. + errChan := make(chan error, 1) + go func() { + for { + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan <- err + return + } + select { + case <-conn.HandshakeComplete(): + conn.CloseWithError(1337, "") + case <-time.After(time.Second): + errChan <- fmt.Errorf("connection did not complete handshake") + } + errChan <- nil + } + }() + + for range 2 { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + require.ErrorIs(t, dial(t, ctx, u), &http3.Error{ErrorCode: 1337, Remote: true}) + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("server did not accept connection") + } + } + } + + // the long request should have been terminated + if graceful { + select { + case err := <-longReqChan: + t.Fatalf("request should not have terminated: %v", err) + case err := <-shutdownChan: + t.Fatalf("graceful shutdown should not have returned: %v", err) + case <-time.After(scaleDuration(10 * time.Millisecond)): + } + + close(handlerChan) + select { + case err := <-longReqChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("long request did not terminate") + } + + select { + case err := <-shutdownChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("shutdown did not complete") + } + } +} diff --git a/third_party/quic-go/integrationtests/self/http_test.go b/third_party/quic-go/integrationtests/self/http_test.go new file mode 100644 index 0000000..30c411b --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_test.go @@ -0,0 +1,1426 @@ +package self_test + +import ( + "bufio" + "bytes" + "compress/gzip" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "maps" + mrand "math/rand/v2" + "net" + "net/http" + "net/http/httptrace" + "net/textproto" + "os" + "strconv" + "strings" + "sync/atomic" + "testing" + "time" + + "golang.org/x/sync/errgroup" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/http3/qlog" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type neverEnding byte + +func (b neverEnding) Read(p []byte) (n int, err error) { + for i := range p { + p[i] = byte(b) + } + return len(p), nil +} + +func randomString(length int) string { + const alphabet = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" + b := make([]byte, length) + for i := range b { + n := mrand.IntN(len(alphabet)) + b[i] = alphabet[n] + } + return string(b) +} + +func startHTTPServer(t *testing.T, mux *http.ServeMux, opts ...func(*http3.Server)) (port int) { + t.Helper() + server := &http3.Server{ + Handler: mux, + TLSConfig: getTLSConfig(), + QUICConfig: getQuicConfig(&quic.Config{Allow0RTT: true, EnableDatagrams: true}), + } + for _, opt := range opts { + opt(server) + } + + conn := newUDPConnLocalhost(t) + done := make(chan struct{}) + go func() { + defer close(done) + server.Serve(conn) + }() + + t.Cleanup(func() { + conn.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("server didn't shut down") + } + }) + return conn.LocalAddr().(*net.UDPAddr).Port +} + +func newHTTP3Client(t *testing.T, opts ...func(*http3.Transport)) *http.Client { + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + QUICConfig: getQuicConfig(&quic.Config{MaxIdleTimeout: 10 * time.Second}), + DisableCompression: true, + } + for _, opt := range opts { + opt(tr) + } + addDialCallback(t, tr) + t.Cleanup(func() { tr.Close() }) + return &http.Client{Transport: tr} +} + +func TestHTTPGet(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/hello", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World!\n") + }) + mux.HandleFunc("/long", func(w http.ResponseWriter, r *http.Request) { + w.Write(PRDataLong) + }) + port := startHTTPServer(t, mux) + + cl := newHTTP3Client(t) + + t.Run("small", func(t *testing.T) { + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 2 * time.Second}) + require.NoError(t, err) + require.Equal(t, "Hello, World!\n", string(body)) + }) + + t.Run("big", func(t *testing.T) { + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/long", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 10 * time.Second}) + require.NoError(t, err) + require.Equal(t, PRDataLong, body) + }) +} + +func TestHTTPPost(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/echo", func(w http.ResponseWriter, r *http.Request) { + io.Copy(w, r.Body) + }) + port := startHTTPServer(t, mux) + + cl := newHTTP3Client(t) + + t.Run("small", func(t *testing.T) { + resp, err := cl.Post( + fmt.Sprintf("https://localhost:%d/echo", port), + "text/plain", + bytes.NewReader([]byte("Hello, world!")), + ) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 2 * time.Second}) + require.NoError(t, err) + require.Equal(t, []byte("Hello, world!"), body) + }) + + t.Run("big", func(t *testing.T) { + resp, err := cl.Post( + fmt.Sprintf("https://localhost:%d/echo", port), + "text/plain", + bytes.NewReader(PRData), + ) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 10 * time.Second}) + require.NoError(t, err) + require.Equal(t, PRData, body) + }) +} + +func TestHTTPMultipleRequests(t *testing.T) { + mux := http.NewServeMux() + port := startHTTPServer(t, mux) + + t.Run("reading the response", func(t *testing.T) { + mux.HandleFunc("/hello", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World!\n") + }) + + cl := newHTTP3Client(t) + var eg errgroup.Group + for range 200 { + eg.Go(func() error { + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", port)) + if err != nil { + return err + } + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("unexpected status code: %d", resp.StatusCode) + } + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 3 * time.Second}) + if err != nil { + return err + } + if string(body) != "Hello, World!\n" { + return fmt.Errorf("unexpected body: %q", body) + } + return nil + }) + } + require.NoError(t, eg.Wait()) + }) + + t.Run("not reading the response", func(t *testing.T) { + mux.HandleFunc("/prdata", func(w http.ResponseWriter, r *http.Request) { + w.Write(PRData) + }) + + cl := newHTTP3Client(t) + const num = 150 + + for range num { + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/prdata", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.NoError(t, resp.Body.Close()) + } + }) +} + +func TestContentLengthForSmallResponse(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/small", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "foo") + io.WriteString(w, "bar") + }) + port := startHTTPServer(t, mux) + + resp, err := newHTTP3Client(t).Get(fmt.Sprintf("https://localhost:%d/small", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.Equal(t, "6", resp.Header.Get("Content-Length")) +} + +func TestHTTPHeaders(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/headers/response", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("foo", "bar") + w.Header().Set("lorem", "ipsum") + w.Header().Set("echo", r.Header.Get("echo")) + }) + port := startHTTPServer(t, mux) + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://localhost:%d/headers/response", port), nil) + require.NoError(t, err) + echoHdr := randomString(128) + req.Header.Set("echo", echoHdr) + + resp, err := newHTTP3Client(t).Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.Equal(t, "bar", resp.Header.Get("foo")) + require.Equal(t, "ipsum", resp.Header.Get("lorem")) + require.Equal(t, echoHdr, resp.Header.Get("echo")) +} + +func TestHTTPHeaderSizeLimitServer(t *testing.T) { + t.Run("large HEADERS frame", func(t *testing.T) { + const limit = 1024 + hdr := make(http.Header) + for range 20 { + hdr.Add(randomString(50), randomString(50)) + } + headersFrameSize := testHTTPHeaderSizeLimitServer(t, hdr, limit) + require.Greater(t, headersFrameSize, limit) + }) + + t.Run("large decompressed HEADERS frame", func(t *testing.T) { + const limit = 1024 + hdr := make(http.Header) + for range 200 { + // This is a QPACK static table entry, so it will be compressed. + hdr.Add("content-type", "text/plain;charset=utf-8") + } + headersFrameSize := testHTTPHeaderSizeLimitServer(t, hdr, limit) + require.Less(t, headersFrameSize, limit) + }) +} + +func testHTTPHeaderSizeLimitServer(t *testing.T, hdr http.Header, limit int) (headersFrameSize int) { + mux := http.NewServeMux() + var handlerCalled bool + mux.HandleFunc("/headers", func(w http.ResponseWriter, r *http.Request) { + handlerCalled = true + }) + port := startHTTPServer(t, mux, func(s *http3.Server) { s.MaxHeaderBytes = limit }) + + var eventRecorder events.Recorder + cl := newHTTP3Client(t, func(tr *http3.Transport) { + tr.QUICConfig = getQuicConfig(&quic.Config{ + MaxIdleTimeout: 10 * time.Second, + Tracer: newTracer(&eventRecorder), + }) + }) + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://localhost:%d/headers", port), nil) + require.NoError(t, err) + req.Header = hdr + + resp, err := cl.Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusRequestHeaderFieldsTooLarge, resp.StatusCode) + require.False(t, handlerCalled) + + for _, ev := range eventRecorder.Events(qlog.FrameCreated{}) { + fc := ev.(qlog.FrameCreated) + if _, ok := fc.Frame.Frame.(qlog.HeadersFrame); ok { + headersFrameSize = fc.Raw.Length + break + } + } + return headersFrameSize +} + +func TestHTTPHeaderSizeLimitClient(t *testing.T) { + t.Run("large HEADERS frame", func(t *testing.T) { + const limit = 1024 + hdr := make(http.Header) + for range 20 { + hdr.Add(randomString(50), randomString(50)) + } + headersFrameSize, requestErr := testHTTPHeaderSizeLimitClient(t, hdr, limit) + require.ErrorContains(t, requestErr, "http3: HEADERS frame too large") + require.Greater(t, headersFrameSize, limit) + }) + + t.Run("large decompressed HEADERS frame", func(t *testing.T) { + const limit = 1024 + hdr := make(http.Header) + for range 200 { + // This is a QPACK static table entry, so it will be compressed. + hdr.Add("content-type", "text/plain;charset=utf-8") + } + headersFrameSize, requestErr := testHTTPHeaderSizeLimitClient(t, hdr, limit) + require.ErrorContains(t, requestErr, "http3: headers too large") + require.Less(t, headersFrameSize, limit) + }) +} + +func testHTTPHeaderSizeLimitClient(t *testing.T, hdr http.Header, limit int) (headersFrameSize int, requestErr error) { + mux := http.NewServeMux() + var handlerCalled atomic.Bool + mux.HandleFunc("/headers", func(w http.ResponseWriter, r *http.Request) { + handlerCalled.Store(true) + for k, v := range hdr { + for _, val := range v { + w.Header().Add(k, val) + } + } + }) + port := startHTTPServer(t, mux) + + var eventRecorder events.Recorder + cl := newHTTP3Client(t, + func(tr *http3.Transport) { + tr.MaxResponseHeaderBytes = limit + tr.QUICConfig = getQuicConfig(&quic.Config{ + MaxIdleTimeout: 10 * time.Second, + Tracer: newTracer(&eventRecorder), + }) + }, + ) + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://localhost:%d/headers", port), nil) + require.NoError(t, err) + + _, requestErr = cl.Do(req) + require.Error(t, requestErr) + require.True(t, handlerCalled.Load()) + + var found bool + for _, ev := range eventRecorder.Events(qlog.FrameParsed{}) { + fp := ev.(qlog.FrameParsed) + if _, ok := fp.Frame.Frame.(qlog.HeadersFrame); ok { + headersFrameSize = fp.Raw.PayloadLength + found = true + break + } + } + require.True(t, found) + return headersFrameSize, requestErr +} + +func TestHTTPResponseTrailers(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/trailers", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Trailer", "AtEnd1, AtEnd2") + w.Header().Add("Trailer", "Never") + w.Header().Add("Trailer", "LAST") + w.Header().Set("Content-Type", "text/plain; charset=utf-8") // normal header + w.WriteHeader(http.StatusOK) + w.Header().Set("AtEnd1", "value 1") + io.WriteString(w, "This HTTP response has both headers before this text and trailers at the end.\n") + w.(http.Flusher).Flush() + w.Header().Set("AtEnd2", "value 2") + io.WriteString(w, "More text\n") + w.(http.Flusher).Flush() + w.Header().Set("LAST", "value 3") + w.Header().Set(http.TrailerPrefix+"Unannounced", "Surprise!") + w.Header().Set("Late-Header", "No surprise!") + }) + + port := startHTTPServer(t, mux) + + resp, err := newHTTP3Client(t).Get(fmt.Sprintf("https://localhost:%d/trailers", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.Empty(t, resp.Header.Get("Trailer")) + require.NotContains(t, resp.Header, "Atend1") + require.NotContains(t, resp.Header, "Atend2") + require.NotContains(t, resp.Header, "Never") + require.NotContains(t, resp.Header, "Last") + require.NotContains(t, resp.Header, "Late-Header") + require.Equal(t, http.Header(map[string][]string{ + "Atend1": nil, + "Atend2": nil, + "Never": nil, + "Last": nil, + }), resp.Trailer) + + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 3 * time.Second}) + require.NoError(t, err) + require.Equal(t, "This HTTP response has both headers before this text and trailers at the end.\nMore text\n", string(body)) + for k := range resp.Header { + require.NotContains(t, k, http.TrailerPrefix) + } + require.Equal(t, http.Header(map[string][]string{ + "Atend1": {"value 1"}, + "Atend2": {"value 2"}, + "Last": {"value 3"}, + "Unannounced": {"Surprise!"}, + }), resp.Trailer) +} + +func TestHTTPRequestTrailers(t *testing.T) { + trailerChan := make(chan http.Header, 2) + bodyChan := make(chan string, 1) + + mux := http.NewServeMux() + mux.HandleFunc("/client-trailers", func(w http.ResponseWriter, r *http.Request) { + trailerBeforeBody := make(http.Header) + maps.Copy(trailerBeforeBody, r.Trailer) + trailerChan <- trailerBeforeBody + + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + bodyChan <- string(body) + + trailer := make(http.Header) + maps.Copy(trailer, r.Trailer) + trailerChan <- trailer + + w.WriteHeader(http.StatusOK) + }) + + port := startHTTPServer(t, mux) + + pr, pw := io.Pipe() + req, err := http.NewRequest(http.MethodPost, fmt.Sprintf("https://localhost:%d/client-trailers", port), pr) + require.NoError(t, err) + req.Trailer = http.Header{ + "Trailer1": nil, + "Trailer2": {"to-be-updated"}, + } + + go func() { + // send the first half of the body + pw.Write(PRData[:len(PRData)/2]) + // then update the trailer values + req.Trailer.Set("Trailer1", "foo") + req.Trailer.Set("Trailer2", "bar") + req.Trailer.Set("Trailer3", "baz") + // send the rest of the body + pw.Write(PRData[len(PRData)/2:]) + pw.Close() + }() + + resp, err := newHTTP3Client(t).Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + select { + case trailersBefore := <-trailerChan: + // trailers before body should have announced keys with nil values + require.Equal(t, http.Header(map[string][]string{"Trailer1": nil, "Trailer2": nil}), trailersBefore) + case <-time.After(time.Second): + t.Fatal("timeout waiting for trailer announcement") + } + + select { + case body := <-bodyChan: + require.Equal(t, string(PRData), body) + case <-time.After(time.Second): + t.Fatal("timeout waiting for body") + } + + select { + case trailers := <-trailerChan: + require.Equal(t, http.Header(map[string][]string{ + "Trailer1": {"foo"}, + "Trailer2": {"bar"}, + "Trailer3": {"baz"}, + }), trailers) + case <-time.After(time.Second): + t.Fatal("timeout waiting for trailers") + } +} + +func TestHTTPErrAbortHandler(t *testing.T) { + respChan := make(chan struct{}) + mux := http.NewServeMux() + mux.HandleFunc("/abort", func(w http.ResponseWriter, r *http.Request) { + // no recover here as it will interfere with the handler + io.WriteString(w, "foobar") + w.(http.Flusher).Flush() + // wait for the client to receive the response + <-respChan + panic(http.ErrAbortHandler) + }) + port := startHTTPServer(t, mux) + + resp, err := newHTTP3Client(t).Get(fmt.Sprintf("https://localhost:%d/abort", port)) + close(respChan) + require.NoError(t, err) + body, err := io.ReadAll(resp.Body) + require.Error(t, err) + var h3Err *http3.Error + require.True(t, errors.As(err, &h3Err)) + require.Equal(t, http3.ErrCodeInternalError, h3Err.ErrorCode) + // the body will be a prefix of what's written + require.True(t, bytes.HasPrefix([]byte("foobar"), body)) +} + +func TestHTTPGzip(t *testing.T) { + mux := http.NewServeMux() + var acceptEncoding string + mux.HandleFunc("/hellogz", func(w http.ResponseWriter, r *http.Request) { + acceptEncoding = r.Header.Get("Accept-Encoding") + w.Header().Set("Content-Encoding", "gzip") + w.Header().Set("foo", "bar") + + gw := gzip.NewWriter(w) + defer gw.Close() + _, err := gw.Write([]byte("Hello, World!\n")) + require.NoError(t, err) + }) + port := startHTTPServer(t, mux) + + cl := newHTTP3Client(t) + cl.Transport.(*http3.Transport).DisableCompression = false + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hellogz", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.True(t, resp.Uncompressed) + + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 3 * time.Second}) + require.NoError(t, err) + require.Equal(t, "Hello, World!\n", string(body)) + + // make sure the server received the Accept-Encoding header + require.Equal(t, "gzip", acceptEncoding) +} + +func TestHTTPDifferentOrigins(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/remote-addr", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-RemoteAddr", r.RemoteAddr) + w.WriteHeader(http.StatusOK) + }) + port := startHTTPServer(t, mux) + + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + QUICConfig: getQuicConfig(nil), + } + t.Cleanup(func() { tr.Close() }) + cl := &http.Client{Transport: tr} + + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/remote-addr", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + addr1 := resp.Header.Get("X-RemoteAddr") + require.NotEmpty(t, addr1) + resp, err = cl.Get(fmt.Sprintf("https://127.0.0.1:%d/remote-addr", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + addr2 := resp.Header.Get("X-RemoteAddr") + require.NotEmpty(t, addr2) + require.Equal(t, addr1, addr2) +} + +func TestHTTPServerIdleTimeout(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/hello", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World!\n") + }) + idleTimeout := scaleDuration(10 * time.Millisecond) + port := startHTTPServer(t, mux, func(s *http3.Server) { s.IdleTimeout = idleTimeout }) + + connChan := make(chan *quic.Conn, 1) + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + Dial: func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + conn, err := quic.DialAddrEarly(ctx, addr, tlsCfg, cfg) + connChan <- conn + return conn, err + }, + } + t.Cleanup(func() { tr.Close() }) + cl := &http.Client{Transport: tr} + + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", port)) + require.NoError(t, err) + // Wait for the server to close the request stream and start the idle timer. + _, err = io.Copy(io.Discard, resp.Body) + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + + var conn *quic.Conn + select { + case conn = <-connChan: + case <-time.After(time.Second): + t.Fatal("connection was not opened") + } + + select { + case <-time.After(3 * idleTimeout): + t.Fatal("connection was not closed") + case <-conn.Context().Done(): + } +} + +func TestHTTPReestablishConnectionAfterDialError(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/hello", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World!\n") + }) + port := startHTTPServer(t, mux) + + var dialCounter int + cl := http.Client{ + Transport: &http3.Transport{ + TLSClientConfig: getTLSClientConfig(), + Dial: func(ctx context.Context, addr string, tlsConf *tls.Config, conf *quic.Config) (*quic.Conn, error) { + dialCounter++ + if dialCounter == 1 { // make the first dial fail + return nil, assert.AnError + } + return quic.DialAddrEarly(ctx, addr, tlsConf, conf) + }, + }, + } + defer cl.Transport.(io.Closer).Close() + + _, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", port)) + require.ErrorIs(t, err, assert.AnError) + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) +} + +func TestHTTPClientRequestContextCancellation(t *testing.T) { + mux := http.NewServeMux() + port := startHTTPServer(t, mux) + cl := newHTTP3Client(t) + + t.Run("before response", func(t *testing.T) { + mux.HandleFunc("/cancel-before", func(w http.ResponseWriter, r *http.Request) { + <-r.Context().Done() + }) + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/cancel-before", port), nil) + require.NoError(t, err) + _, err = cl.Do(req) + require.Error(t, err) + require.ErrorIs(t, err, context.DeadlineExceeded) + }) + + t.Run("after response", func(t *testing.T) { + errChan := make(chan error, 1) + mux.HandleFunc("/cancel-after", func(w http.ResponseWriter, r *http.Request) { + // TODO(#4508): check for request context cancellations + for { + if _, err := io.WriteString(w, "foobar"); err != nil { + errChan <- err + return + } + } + }) + + ctx, cancel := context.WithCancel(context.Background()) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/cancel-after", port), nil) + require.NoError(t, err) + resp, err := cl.Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + cancel() + + select { + case err := <-errChan: + require.Error(t, err) + var http3Err *http3.Error + require.True(t, errors.As(err, &http3Err)) + require.Equal(t, http3.ErrCodeRequestCanceled, http3Err.ErrorCode) + require.True(t, http3Err.Remote) + case <-time.After(time.Second): + t.Fatal("handler was not called") + } + + _, err = resp.Body.Read([]byte{0}) + var http3Err *http3.Error + require.True(t, errors.As(err, &http3Err)) + require.Equal(t, http3.ErrCodeRequestCanceled, http3Err.ErrorCode) + require.False(t, http3Err.Remote) + }) +} + +func TestHTTPDeadlines(t *testing.T) { + const deadlineDelay = 50 * time.Millisecond + + mux := http.NewServeMux() + port := startHTTPServer(t, mux) + cl := newHTTP3Client(t) + + t.Run("read deadline", func(t *testing.T) { + type result struct { + body []byte + err error + } + + resultChan := make(chan result, 1) + mux.HandleFunc("/read-deadline", func(w http.ResponseWriter, r *http.Request) { + rc := http.NewResponseController(w) + require.NoError(t, rc.SetReadDeadline(time.Now().Add(deadlineDelay))) + body, err := io.ReadAll(r.Body) + resultChan <- result{body: body, err: err} + io.WriteString(w, "ok") + }) + + expectedEnd := time.Now().Add(deadlineDelay) + resp, err := cl.Post( + fmt.Sprintf("https://localhost:%d/read-deadline", port), + "text/plain", + neverEnding('a'), + ) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 2 * deadlineDelay}) + require.NoError(t, err) + require.True(t, time.Now().After(expectedEnd)) + require.Equal(t, "ok", string(body)) + + select { + case result := <-resultChan: + require.ErrorIs(t, result.err, os.ErrDeadlineExceeded) + require.Contains(t, string(result.body), "aa") + default: + t.Fatal("handler was not called") + } + }) + + t.Run("write deadline", func(t *testing.T) { + errChan := make(chan error, 1) + mux.HandleFunc("/write-deadline", func(w http.ResponseWriter, r *http.Request) { + rc := http.NewResponseController(w) + require.NoError(t, rc.SetWriteDeadline(time.Now().Add(deadlineDelay))) + + _, err := io.Copy(w, neverEnding('a')) + errChan <- err + }) + + expectedEnd := time.Now().Add(deadlineDelay) + + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/write-deadline", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + body, err := io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 2 * deadlineDelay}) + require.NoError(t, err) + require.True(t, time.Now().After(expectedEnd)) + require.Contains(t, string(body), "aa") + + select { + case err := <-errChan: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + case <-time.After(2 * deadlineDelay): + t.Fatal("handler was not called") + } + }) +} + +func TestHTTPServeQUICConn(t *testing.T) { + tlsConf := getTLSConfig() + tlsConf.NextProtos = []string{http3.NextProtoH3} + ln, err := quic.Listen(newUDPConnLocalhost(t), tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + mux := http.NewServeMux() + mux.HandleFunc("/hello", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, "Hello, World!\n") + }) + server := &http3.Server{ + TLSConfig: tlsConf, + QUICConfig: getQuicConfig(nil), + Handler: mux, + } + errChan := make(chan error, 1) + go func() { + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan <- fmt.Errorf("failed to accept QUIC connection: %w", err) + return + } + errChan <- server.ServeQUICConn(conn) // returns once the client closes + }() + + cl := newHTTP3Client(t) + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/hello", ln.Addr().(*net.UDPAddr).Port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + require.NoError(t, cl.Transport.(io.Closer).Close()) + select { + case err := <-errChan: + require.Error(t, err) + require.ErrorContains(t, err, "accepting stream failed") + case <-time.After(time.Second): + t.Fatal("server didn't shut down") + } +} + +func TestHTTPContextFromQUIC(t *testing.T) { + conn := newUDPConnLocalhost(t) + tr := &quic.Transport{ + Conn: conn, + ConnContext: func(ctx context.Context, _ *quic.ClientInfo) (context.Context, error) { + return context.WithValue(ctx, "foo", "bar"), nil + }, + } + defer tr.Close() + tlsConf := getTLSConfig() + tlsConf.NextProtos = []string{http3.NextProtoH3} + ln, err := tr.Listen(tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + mux := http.NewServeMux() + ctxChan := make(chan context.Context, 1) + mux.HandleFunc("/quic-conn-context", func(w http.ResponseWriter, r *http.Request) { + ctxChan <- r.Context() + }) + + server := &http3.Server{Handler: mux} + go func() { + c, err := ln.Accept(context.Background()) + require.NoError(t, err) + server.ServeQUICConn(c) + }() + + cl := newHTTP3Client(t) + resp, err := cl.Get(fmt.Sprintf("https://localhost:%d/quic-conn-context", conn.LocalAddr().(*net.UDPAddr).Port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + select { + case ctx := <-ctxChan: + v, ok := ctx.Value("foo").(string) + require.True(t, ok) + require.Equal(t, "bar", v) + default: + t.Fatal("context not set") + } +} + +func TestHTTPConnContext(t *testing.T) { + mux := http.NewServeMux() + requestCtxChan := make(chan context.Context, 1) + mux.HandleFunc("/context", func(w http.ResponseWriter, r *http.Request) { + requestCtxChan <- r.Context() + }) + + var server *http3.Server + connCtxChan := make(chan context.Context, 1) + port := startHTTPServer(t, + mux, + func(s *http3.Server) { server = s }, + func(s *http3.Server) { + s.ConnContext = func(ctx context.Context, c *quic.Conn) context.Context { + connCtxChan <- ctx + ctx = context.WithValue(ctx, "foo", "bar") + return ctx + } + }, + ) + + resp, err := newHTTP3Client(t).Get(fmt.Sprintf("https://localhost:%d/context", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + select { + case ctx := <-connCtxChan: + serv, ok := ctx.Value(http3.ServerContextKey).(*http3.Server) + require.True(t, ok) + require.Equal(t, server, serv) + default: + t.Fatal("handler was not called") + } + + select { + case ctx := <-requestCtxChan: + v, ok := ctx.Value("foo").(string) + require.True(t, ok) + require.Equal(t, "bar", v) + + serv, ok := ctx.Value(http3.ServerContextKey).(*http3.Server) + require.True(t, ok) + require.Equal(t, server, serv) + default: + t.Fatal("handler was not called") + } +} + +func TestHTTPRemoteAddrContextKey(t *testing.T) { + ctxChan := make(chan context.Context, 1) + mux := http.NewServeMux() + mux.HandleFunc("/remote-addr", func(w http.ResponseWriter, r *http.Request) { + ctxChan <- r.Context() + }) + + port := startHTTPServer(t, mux) + + resp, err := newHTTP3Client(t).Get(fmt.Sprintf("https://localhost:%d/remote-addr", port)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + select { + case ctx := <-ctxChan: + _, ok := ctx.Value(http3.RemoteAddrContextKey).(net.Addr) + require.True(t, ok) + require.Equal(t, "127.0.0.1", ctx.Value(http3.RemoteAddrContextKey).(*net.UDPAddr).IP.String()) + default: + t.Fatal("handler was not called") + } +} + +func TestHTTPStreamedRequests(t *testing.T) { + errChan := make(chan error, 1) + mux := http.NewServeMux() + mux.HandleFunc("/echoline", func(w http.ResponseWriter, r *http.Request) { + defer close(errChan) + w.WriteHeader(200) + w.(http.Flusher).Flush() + reader := bufio.NewReader(r.Body) + for { + msg, err := reader.ReadString('\n') + if err != nil { + return + } + if _, err := io.WriteString(w, msg); err != nil { + errChan <- err + return + } + w.(http.Flusher).Flush() + } + }) + + port := startHTTPServer(t, mux) + + r, w := io.Pipe() + req, err := http.NewRequest(http.MethodPut, fmt.Sprintf("https://localhost:%d/echoline", port), r) + require.NoError(t, err) + client := newHTTP3Client(t) + rsp, err := client.Do(req) + require.NoError(t, err) + require.Equal(t, 200, rsp.StatusCode) + + reader := bufio.NewReader(rsp.Body) + for i := range 5 { + msg := fmt.Sprintf("Hello world, %d!\n", i) + fmt.Fprint(w, msg) + msgRcvd, err := reader.ReadString('\n') + require.NoError(t, err) + require.Equal(t, msg, msgRcvd) + } + require.NoError(t, req.Body.Close()) + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("handler did not complete") + } +} + +func TestHTTP1xxResponse(t *testing.T) { + header1 := "; rel=preload; as=style" + header2 := "; rel=preload; as=script" + data := "1xx-test-data" + mux := http.NewServeMux() + mux.HandleFunc("/103-early-data", func(w http.ResponseWriter, r *http.Request) { + w.Header().Add("Link", header1) + w.Header().Add("Link", header2) + w.WriteHeader(http.StatusEarlyHints) + io.WriteString(w, data) + w.WriteHeader(http.StatusOK) + }) + + port := startHTTPServer(t, mux) + + var ( + cnt int + status int + hdr textproto.MIMEHeader + ) + ctx := httptrace.WithClientTrace(context.Background(), &httptrace.ClientTrace{ + Got1xxResponse: func(code int, header textproto.MIMEHeader) error { + hdr = header + status = code + cnt++ + return nil + }, + }) + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/103-early-data", port), nil) + require.NoError(t, err) + resp, err := newHTTP3Client(t).Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, data, string(body)) + require.Equal(t, http.StatusEarlyHints, status) + require.Equal(t, []string{header1, header2}, hdr.Values("Link")) + require.Equal(t, 1, cnt) + require.Equal(t, []string{header1, header2}, resp.Header.Values("Link")) + require.NoError(t, resp.Body.Close()) +} + +func TestHTTP1xxTerminalResponse(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/101-switch-protocols", func(w http.ResponseWriter, r *http.Request) { + w.Header().Add("foo", "bar") + w.WriteHeader(http.StatusSwitchingProtocols) + }) + + port := startHTTPServer(t, mux) + + var ( + cnt int + status int + ) + ctx := httptrace.WithClientTrace(context.Background(), &httptrace.ClientTrace{ + Got1xxResponse: func(code int, header textproto.MIMEHeader) error { + status = code + cnt++ + return nil + }, + }) + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/101-switch-protocols", port), nil) + require.NoError(t, err) + resp, err := newHTTP3Client(t).Do(req) + require.NoError(t, err) + require.Equal(t, http.StatusSwitchingProtocols, resp.StatusCode) + require.Equal(t, "bar", resp.Header.Get("Foo")) + require.Zero(t, status) + require.Zero(t, cnt) + require.NoError(t, resp.Body.Close()) +} + +func TestHTTP0RTT(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/0rtt", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, strconv.FormatBool(!r.TLS.HandshakeComplete)) + }) + port := startHTTPServer(t, mux) + + var num0RTTPackets atomic.Uint32 + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: port}, + DelayPacket: func(_ quicproxy.Direction, _, _ net.Addr, data []byte) time.Duration { + if containsPacketType(data, protocol.PacketType0RTT) { + num0RTTPackets.Add(1) + } + return scaleDuration(25 * time.Millisecond) + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + tlsConf := getTLSClientConfigWithoutServerName() + puts := make(chan string, 10) + tlsConf.ClientSessionCache = newClientSessionCache(tls.NewLRUClientSessionCache(10), nil, puts) + tr := &http3.Transport{ + TLSClientConfig: tlsConf, + QUICConfig: getQuicConfig(&quic.Config{MaxIdleTimeout: 10 * time.Second}), + DisableCompression: true, + } + defer tr.Close() + addDialCallback(t, tr) + + proxyPort := proxy.LocalAddr().(*net.UDPAddr).Port + req, err := http.NewRequest(http3.MethodGet0RTT, fmt.Sprintf("https://localhost:%d/0rtt", proxyPort), nil) + require.NoError(t, err) + rsp, err := tr.RoundTrip(req) + require.NoError(t, err) + require.Equal(t, 200, rsp.StatusCode) + data, err := io.ReadAll(rsp.Body) + require.NoError(t, err) + require.Equal(t, "false", string(data)) + require.Zero(t, num0RTTPackets.Load()) + + select { + case <-puts: + case <-time.After(time.Second): + t.Fatal("did not receive session ticket") + } + + tr2 := &http3.Transport{ + TLSClientConfig: tr.TLSClientConfig, + QUICConfig: tr.QUICConfig, + DisableCompression: true, + } + defer tr2.Close() + addDialCallback(t, tr2) + rsp, err = tr2.RoundTrip(req) + require.NoError(t, err) + require.Equal(t, 200, rsp.StatusCode) + data, err = io.ReadAll(rsp.Body) + require.NoError(t, err) + require.Equal(t, "true", string(data)) + require.NotZero(t, num0RTTPackets.Load()) +} + +func TestHTTPStreamer(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/httpstreamer", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + + str := w.(http3.HTTPStreamer).HTTPStream() + str.Write([]byte("foobar")) + + // Do this in a Go routine, so that the handler returns early. + // This way, we can also check that the HTTP/3 doesn't close the stream. + go func() { + defer str.Close() + _, _ = io.Copy(str, str) + }() + }) + + port := startHTTPServer(t, mux) + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("https://localhost:%d/httpstreamer", port), nil) + require.NoError(t, err) + tlsConf := getTLSClientConfig() + tlsConf.NextProtos = []string{http3.NextProtoH3} + ctx := t.Context() + conn, err := quic.Dial(ctx, newUDPConnLocalhost(t), &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: port}, tlsConf, getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + tr := http3.Transport{} + addDialCallback(t, &tr) + cc := tr.NewClientConn(conn) + str, err := cc.OpenRequestStream(ctx) + require.NoError(t, err) + require.NoError(t, str.SendRequestHeader(req)) + + rsp, err := str.ReadResponse() + require.NoError(t, err) + require.Equal(t, 200, rsp.StatusCode) + + b := make([]byte, 6) + _, err = io.ReadFull(str, b) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), b) + + _, err = str.Write(PRData) + require.NoError(t, err) + require.NoError(t, str.Close()) + repl, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRData, repl) +} + +type blackHoleConn struct { + net.PacketConn + block atomic.Bool + close chan struct{} +} + +func (c *blackHoleConn) WriteTo(b []byte, addr net.Addr) (int, error) { + return c.PacketConn.WriteTo(b, addr) +} + +func (c *blackHoleConn) ReadFrom(b []byte) (int, net.Addr, error) { + if c.block.Load() { + <-c.close + return 0, nil, errors.New("blocked") + } + n, _, err := c.PacketConn.ReadFrom(b) + if c.block.Load() { + <-c.close + return 0, nil, errors.New("blocked") + } + return n, nil, err +} + +func (c *blackHoleConn) Close() error { + close(c.close) + return c.PacketConn.Close() +} + +func (c *blackHoleConn) StartBlocking() { c.block.Store(true) } + +func TestHTTPRequestRetryAfterIdleTimeout(t *testing.T) { + t.Run("only cached conn", func(t *testing.T) { + testHTTPRequestRetryAfterIdleTimeout(t, true) + }) + t.Run("allow re-dialing", func(t *testing.T) { + testHTTPRequestRetryAfterIdleTimeout(t, false) + }) +} + +func testHTTPRequestRetryAfterIdleTimeout(t *testing.T, onlyCachedConn bool) { + t.Setenv("QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING", "true") + + mux := http.NewServeMux() + mux.HandleFunc("/remote-addr", func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, r.RemoteAddr) + }) + port := startHTTPServer(t, mux, func(s *http3.Server) {}) + + firstConn := &blackHoleConn{PacketConn: newUDPConnLocalhost(t), close: make(chan struct{})} + secondConn := newUDPConnLocalhost(t) + conns := []net.PacketConn{firstConn, secondConn} + require.NotEqual(t, firstConn.LocalAddr().String(), secondConn.LocalAddr().String()) + + idleTimeout := scaleDuration(10 * time.Millisecond) + connChan := make(chan *quic.Conn, 2) + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + QUICConfig: getQuicConfig(&quic.Config{MaxIdleTimeout: idleTimeout}), + Dial: func(ctx context.Context, a string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + conn := conns[0] + conns = conns[1:] + addr, err := net.ResolveUDPAddr("udp", a) + if err != nil { + return nil, err + } + c, err := quic.DialEarly(ctx, conn, addr, tlsCfg, cfg) + if err != nil { + return nil, err + } + connChan <- c + return c, nil + }, + DisableCompression: true, + } + t.Cleanup(func() { tr.Close() }) + + var headersCount int + req, err := http.NewRequestWithContext( + httptrace.WithClientTrace(context.Background(), &httptrace.ClientTrace{ + WroteHeaders: func() { headersCount++ }, + }), + http.MethodGet, + fmt.Sprintf("https://127.0.0.1:%d/remote-addr", port), + // Add a body (wrappped so that http.NewRequest doesn't set the GetBody callback), + // to make it impossible to retry this request. + // This tests that the detection logic works properly: + // If the request fails before the stream can be opened, it is always safe to retry. + io.LimitReader(strings.NewReader("foobar"), 1000), + ) + require.NoError(t, err) + + resp, err := tr.RoundTripOpt(req, http3.RoundTripOpt{}) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, firstConn.LocalAddr().String(), string(body)) + + firstConn.StartBlocking() + // wait for the connection to time out + select { + case c := <-connChan: + select { + case <-c.Context().Done(): + case <-time.After(time.Second): + t.Fatal("connection did not time out") + } + case <-time.After(time.Second): + t.Fatal("no connection was created") + } + + // second request should succeed after re-dialing + resp, err = tr.RoundTripOpt(req, http3.RoundTripOpt{OnlyCachedConn: onlyCachedConn}) + if onlyCachedConn { + require.EqualError(t, err, "http3: no cached connection was available") + require.Len(t, conns, 1) // no second dial attempt + require.Equal(t, 1, headersCount) + return + } + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err = io.ReadAll(&readerWithTimeout{Reader: resp.Body, Timeout: 2 * time.Second}) + require.NoError(t, err) + require.Equal(t, secondConn.LocalAddr().String(), string(body)) + + require.Equal(t, 2, headersCount) + require.Empty(t, conns) // make sure we dialed 2 connections +} + +func TestHTTPRequestAfterGracefulShutdown(t *testing.T) { + t.Run("Request.GetBody set", func(t *testing.T) { + testHTTPRequestAfterGracefulShutdown(t, true) + }) + t.Run("Request.GetBody not set", func(t *testing.T) { + testHTTPRequestAfterGracefulShutdown(t, false) + }) +} + +func testHTTPRequestAfterGracefulShutdown(t *testing.T, setGetBody bool) { + t.Setenv("QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING", "true") + + ln, err := quic.ListenEarly( + newUDPConnLocalhost(t), + http3.ConfigureTLSConfig(getTLSConfig()), + getQuicConfig(nil), + ) + require.NoError(t, err) + + var inShutdown atomic.Bool + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(_ quicproxy.Direction, _, _ net.Addr, data []byte) time.Duration { + if inShutdown.Load() { + return scaleDuration(10 * time.Millisecond) + } + return scaleDuration(2 * time.Millisecond) + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + mux2 := http.NewServeMux() + mux2.HandleFunc("/echo", func(w http.ResponseWriter, r *http.Request) { + data, _ := io.ReadAll(r.Body) + w.Write(data) + }) + server2 := &http3.Server{Handler: mux2} + + done := make(chan struct{}) + defer close(done) + server1 := &http3.Server{Handler: http.NewServeMux()} + + go server1.ServeListener(ln) + + tlsConf := getTLSClientConfigWithoutServerName() + tlsConf.NextProtos = []string{http3.NextProtoH3} + var dialCount int + tr := &http3.Transport{ + TLSClientConfig: tlsConf, + Dial: func(ctx context.Context, a string, tlsConf *tls.Config, conf *quic.Config) (*quic.Conn, error) { + addr, err := net.ResolveUDPAddr("udp", a) + if err != nil { + return nil, err + } + dialCount++ + return quic.DialEarly(ctx, newUDPConnLocalhost(t), addr, tlsConf, conf) + }, + } + t.Cleanup(func() { tr.Close() }) + cl := &http.Client{Transport: tr} + + // first request to establish the connection + resp, err := cl.Get(fmt.Sprintf("https://%s/", proxy.LocalAddr())) + require.NoError(t, err) + require.Equal(t, http.StatusNotFound, resp.StatusCode) + + // If the body is a strings.Reader, http.NewRequest automatically sets the GetBody callback. + // This can be prevented by using a different kind of reader, e.g. the io.LimitReader. + var headersCount int + req, err := http.NewRequestWithContext( + httptrace.WithClientTrace(context.Background(), &httptrace.ClientTrace{ + WroteHeaders: func() { headersCount++ }, + }), + http.MethodGet, + fmt.Sprintf("https://%s/echo", proxy.LocalAddr()), + io.LimitReader(strings.NewReader("foobar"), 1000), + ) + require.NoError(t, err) + if setGetBody { + req.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader("foobaz")), nil + } + } else { + require.Nil(t, req.GetBody) + } + + // By increasing the RTT, we make sure that the request is sent before the client receives the GOAWAY frame. + inShutdown.Store(true) + go server1.Shutdown(context.Background()) + go server2.ServeListener(ln) + defer server2.Close() + + resp, err = cl.Do(req) + if !setGetBody { + require.ErrorContains(t, err, "after Request.Body was written; define Request.GetBody to avoid this error") + require.Equal(t, 1, dialCount) + require.Equal(t, 1, headersCount) + return + } + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, "foobaz", string(body)) + require.Equal(t, 2, dialCount) + require.Equal(t, 2, headersCount) +} diff --git a/third_party/quic-go/integrationtests/self/http_trace_test.go b/third_party/quic-go/integrationtests/self/http_trace_test.go new file mode 100644 index 0000000..77fc35c --- /dev/null +++ b/third_party/quic-go/integrationtests/self/http_trace_test.go @@ -0,0 +1,137 @@ +package self_test + +import ( + "context" + "crypto/tls" + "fmt" + "net" + "net/http" + "net/http/httptrace" + "net/textproto" + "testing" + "time" + + "github.com/apernet/quic-go/http3" + "github.com/stretchr/testify/require" +) + +func TestHTTPClientTrace(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/client-trace", func(w http.ResponseWriter, r *http.Request) { + time.Sleep(100 * time.Millisecond) + w.WriteHeader(http.StatusContinue) + }) + port := startHTTPServer(t, mux) + + buf := make([]byte, 1) + type event struct { + Key string + Args any + } + eventQueue := make(chan event, 100) + wait100Continue := false + trace := httptrace.ClientTrace{ + GetConn: func(hostPort string) { eventQueue <- event{Key: "GetConn", Args: hostPort} }, + GotConn: func(info httptrace.GotConnInfo) { eventQueue <- event{Key: "GotConn", Args: info} }, + GotFirstResponseByte: func() { eventQueue <- event{Key: "GotFirstResponseByte"} }, + Got100Continue: func() { eventQueue <- event{Key: "Got100Continue"} }, + Got1xxResponse: func(code int, header textproto.MIMEHeader) error { + eventQueue <- event{Key: "Got1xxResponse", Args: code} + return nil + }, + DNSStart: func(di httptrace.DNSStartInfo) { eventQueue <- event{Key: "DNSStart", Args: di} }, + DNSDone: func(di httptrace.DNSDoneInfo) { eventQueue <- event{Key: "DNSDone", Args: di} }, + ConnectStart: func(network, addr string) { + eventQueue <- event{Key: "ConnectStart", Args: map[string]string{"network": network, "addr": addr}} + }, + ConnectDone: func(network, addr string, err error) { + eventQueue <- event{Key: "ConnectDone", Args: map[string]any{"network": network, "addr": addr, "err": err}} + }, + TLSHandshakeStart: func() { eventQueue <- event{Key: "TLSHandshakeStart"} }, + TLSHandshakeDone: func(state tls.ConnectionState, err error) { + eventQueue <- event{Key: "TLSHandshakeDone", Args: map[string]any{"state": state, "err": err}} + }, + WroteHeaderField: func(key string, value []string) { + if key != ":authority" { + return + } + eventQueue <- event{Key: "WroteHeaderField", Args: value[0]} + }, + WroteHeaders: func() { eventQueue <- event{Key: "WroteHeaders"} }, + Wait100Continue: func() { wait100Continue = true }, + WroteRequest: func(i httptrace.WroteRequestInfo) { eventQueue <- event{Key: "WroteRequest", Args: i} }, + } + ctx := httptrace.WithClientTrace(context.Background(), &trace) + + tr := &http3.Transport{ + TLSClientConfig: getTLSClientConfigWithoutServerName(), + QUICConfig: getQuicConfig(nil), + } + t.Cleanup(func() { tr.Close() }) + cl := &http.Client{Transport: tr} + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("https://localhost:%d/client-trace", port), nil) + require.NoError(t, err) + resp, err := cl.Do(req) + close(eventQueue) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + events := make([]string, 0, len(eventQueue)) + for e := range eventQueue { + events = append(events, e.Key) + switch e.Key { + case "GetConn": + require.Equal(t, fmt.Sprintf("localhost:%d", port), e.Args.(string)) + case "GotConn": + info := e.Args.(httptrace.GotConnInfo) + require.Equal(t, fmt.Sprintf("127.0.0.1:%d", port), info.Conn.RemoteAddr().String()) + host, _, err := net.SplitHostPort(info.Conn.LocalAddr().String()) + require.NoError(t, err) + require.Contains(t, []string{"::", "0.0.0.0"}, host) + require.Panics(t, func() { info.Conn.Close() }) + require.Panics(t, func() { info.Conn.Read(buf) }) + require.Panics(t, func() { info.Conn.Write(buf) }) + require.Panics(t, func() { info.Conn.SetDeadline(time.Now()) }) + require.Panics(t, func() { info.Conn.SetReadDeadline(time.Now()) }) + require.Panics(t, func() { info.Conn.SetWriteDeadline(time.Now()) }) + case "Got1xxResponse": + require.Equal(t, 100, e.Args.(int)) + case "DNSStart": + require.Equal(t, "localhost", e.Args.(httptrace.DNSStartInfo).Host) + case "DNSDone": + require.Condition(t, func() bool { + localhost := net.IPv4(127, 0, 0, 1) + localhostTo16 := localhost.To16() + for _, addr := range e.Args.(httptrace.DNSDoneInfo).Addrs { + if addr.IP.Equal(localhost) || addr.IP.Equal(localhostTo16) { + return true + } + } + return false + }) + case "ConnectStart": + require.Equal(t, "udp", e.Args.(map[string]string)["network"]) + require.Equal(t, fmt.Sprintf("127.0.0.1:%d", port), e.Args.(map[string]string)["addr"]) + case "ConnectDone": + require.Equal(t, "udp", e.Args.(map[string]any)["network"]) + require.Equal(t, fmt.Sprintf("127.0.0.1:%d", port), e.Args.(map[string]any)["addr"]) + require.Nil(t, e.Args.(map[string]any)["err"]) + case "TLSHandshakeDone": + require.Nil(t, e.Args.(map[string]any)["err"]) + state := e.Args.(map[string]any)["state"].(tls.ConnectionState) + require.Equal(t, 1, len(state.PeerCertificates)) + require.Equal(t, "localhost", state.PeerCertificates[0].DNSNames[0]) + case "WroteHeaderField": + require.Equal(t, fmt.Sprintf("localhost:%d", port), e.Args.(string)) + case "WroteRequest": + require.NoError(t, e.Args.(httptrace.WroteRequestInfo).Err) + } + } + require.Equal(t, + []string{ + "GetConn", "DNSStart", "DNSDone", "ConnectStart", "TLSHandshakeStart", "TLSHandshakeDone", + "ConnectDone", "GotConn", "WroteHeaderField", "WroteHeaders", "WroteRequest", + "GotFirstResponseByte", "Got1xxResponse", "Got100Continue", + }, events) + require.Falsef(t, wait100Continue, "wait 100 continue") // Note: not supported Expect: 100-continue +} diff --git a/third_party/quic-go/integrationtests/self/key_update_test.go b/third_party/quic-go/integrationtests/self/key_update_test.go new file mode 100644 index 0000000..47df4a9 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/key_update_test.go @@ -0,0 +1,93 @@ +package self_test + +import ( + "context" + "io" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestKeyUpdates(t *testing.T) { + reset := handshake.SetKeyUpdateInterval(1) // update keys as frequently as possible + t.Cleanup(reset) + + countKeyPhases := func(events []qlogwriter.Event) (sent, received int) { + lastKeyPhaseSend := protocol.KeyPhaseOne + lastKeyPhaseReceive := protocol.KeyPhaseOne + for _, ev := range events { + switch ev := ev.(type) { + case qlog.PacketSent: + if ev.Header.KeyPhaseBit != lastKeyPhaseSend { + sent++ + lastKeyPhaseSend = ev.Header.KeyPhaseBit + } + case qlog.PacketReceived: + if ev.Header.KeyPhaseBit != lastKeyPhaseReceive { + received++ + lastKeyPhaseReceive = ev.Header.KeyPhaseBit + } + } + } + return + } + + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), nil) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + var eventRecorder events.Recorder + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{Tracer: newTracer(&eventRecorder)}), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + serverErrChan := make(chan error, 1) + go func() { + str, err := serverConn.OpenUniStream() + if err != nil { + serverErrChan <- err + return + } + defer str.Close() + if _, err := str.Write(PRDataLong); err != nil { + serverErrChan <- err + return + } + close(serverErrChan) + }() + + str, err := conn.AcceptUniStream(ctx) + require.NoError(t, err) + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRDataLong, data) + require.NoError(t, conn.CloseWithError(0, "")) + + require.NoError(t, <-serverErrChan) + + keyPhasesSent, keyPhasesReceived := countKeyPhases(eventRecorder.Events()) + t.Logf("Used %d key phases on outgoing and %d key phases on incoming packets.", keyPhasesSent, keyPhasesReceived) + assert.Greater(t, keyPhasesReceived, 10) + assert.InDelta(t, keyPhasesSent, keyPhasesReceived, 2) +} diff --git a/third_party/quic-go/integrationtests/self/mitm_test.go b/third_party/quic-go/integrationtests/self/mitm_test.go new file mode 100644 index 0000000..c7722a2 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/mitm_test.go @@ -0,0 +1,428 @@ +package self_test + +import ( + "context" + "crypto/rand" + "errors" + "io" + "math" + mrand "math/rand/v2" + "net" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/testutils" + + "github.com/stretchr/testify/require" +) + +const mitmTestConnIDLen = 6 + +func getTransportsForMITMTest(t *testing.T) (serverTransport, clientTransport *quic.Transport) { + serverTransport = &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnectionIDLength: mitmTestConnIDLen, + } + addTracer(serverTransport) + t.Cleanup(func() { serverTransport.Close() }) + + clientTransport = &quic.Transport{ + Conn: newUDPConnLocalhost(t), + ConnectionIDLength: mitmTestConnIDLen, + } + addTracer(clientTransport) + t.Cleanup(func() { clientTransport.Close() }) + + return serverTransport, clientTransport +} + +func TestMITMInjectRandomPackets(t *testing.T) { + t.Run("towards the server", func(t *testing.T) { + testMITMInjectRandomPackets(t, quicproxy.DirectionIncoming) + }) + + t.Run("towards the client", func(t *testing.T) { + testMITMInjectRandomPackets(t, quicproxy.DirectionOutgoing) + }) +} + +func TestMITMDuplicatePackets(t *testing.T) { + t.Run("towards the server", func(t *testing.T) { + testMITMDuplicatePackets(t, quicproxy.DirectionIncoming) + }) + + t.Run("towards the client", func(t *testing.T) { + testMITMDuplicatePackets(t, quicproxy.DirectionOutgoing) + }) +} + +func TestMITCorruptPackets(t *testing.T) { + t.Run("towards the server", func(t *testing.T) { + testMITMCorruptPackets(t, quicproxy.DirectionIncoming) + }) + + t.Run("towards the client", func(t *testing.T) { + testMITMCorruptPackets(t, quicproxy.DirectionOutgoing) + }) +} + +func testMITMInjectRandomPackets(t *testing.T, direction quicproxy.Direction) { + createRandomPacketOfSameType := func(b []byte) []byte { + if wire.IsLongHeaderPacket(b[0]) { + hdr, _, _, err := wire.ParsePacket(b) + if err != nil { + return nil + } + replyHdr := &wire.ExtendedHeader{ + Header: wire.Header{ + DestConnectionID: hdr.DestConnectionID, + SrcConnectionID: hdr.SrcConnectionID, + Type: hdr.Type, + Version: hdr.Version, + }, + PacketNumber: protocol.PacketNumber(mrand.Int32N(math.MaxInt32 / 4)), + PacketNumberLen: protocol.PacketNumberLen(mrand.IntN(4) + 1), + } + payloadLen := mrand.IntN(100) + replyHdr.Length = protocol.ByteCount(mrand.IntN(payloadLen + 1)) + data, err := replyHdr.Append(nil, hdr.Version) + if err != nil { + panic("failed to append header: " + err.Error()) + } + r := make([]byte, payloadLen) + rand.Read(r) + return append(data, r...) + } + // short header packet + connID, err := wire.ParseConnectionID(b, mitmTestConnIDLen) + if err != nil { + return nil + } + _, pn, pnLen, _, err := wire.ParseShortHeader(b, mitmTestConnIDLen) + if err != nil && !errors.Is(err, wire.ErrInvalidReservedBits) { // normally, ParseShortHeader is called after decrypting the header + panic("failed to parse short header: " + err.Error()) + } + data, err := wire.AppendShortHeader(nil, connID, pn, pnLen, protocol.KeyPhaseBit(mrand.IntN(2))) + if err != nil { + return nil + } + payloadLen := mrand.IntN(100) + r := make([]byte, payloadLen) + rand.Read(r) + return append(data, r...) + } + + rtt := scaleDuration(10 * time.Millisecond) + serverTransport, clientTransport := getTransportsForMITMTest(t) + + dropCallback := func(dir quicproxy.Direction, _, _ net.Addr, b []byte) bool { + if dir != direction { + return false + } + go func() { + ticker := time.NewTicker(rtt / 10) + defer ticker.Stop() + for range 10 { + switch direction { + case quicproxy.DirectionIncoming: + clientTransport.WriteTo(createRandomPacketOfSameType(b), serverTransport.Conn.LocalAddr()) + case quicproxy.DirectionOutgoing: + serverTransport.WriteTo(createRandomPacketOfSameType(b), clientTransport.Conn.LocalAddr()) + } + <-ticker.C + } + }() + return false + } + + runMITMTest(t, serverTransport, clientTransport, rtt, dropCallback) +} + +func testMITMDuplicatePackets(t *testing.T, direction quicproxy.Direction) { + serverTransport, clientTransport := getTransportsForMITMTest(t) + rtt := scaleDuration(10 * time.Millisecond) + + dropCallback := func(dir quicproxy.Direction, _, _ net.Addr, b []byte) bool { + if dir != direction { + return false + } + switch direction { + case quicproxy.DirectionIncoming: + clientTransport.WriteTo(b, serverTransport.Conn.LocalAddr()) + case quicproxy.DirectionOutgoing: + serverTransport.WriteTo(b, clientTransport.Conn.LocalAddr()) + } + return false + } + + runMITMTest(t, serverTransport, clientTransport, rtt, dropCallback) +} + +func testMITMCorruptPackets(t *testing.T, direction quicproxy.Direction) { + serverTransport, clientTransport := getTransportsForMITMTest(t) + rtt := scaleDuration(5 * time.Millisecond) + + var numCorrupted atomic.Int32 + dropCallback := func(dir quicproxy.Direction, _, _ net.Addr, b []byte) bool { + if dir != direction { + return false + } + isLongHeaderPacket := wire.IsLongHeaderPacket(b[0]) + // corrupt 20% of long header packets and 5% of short header packets + if isLongHeaderPacket && mrand.IntN(4) != 0 { + return false + } + if !isLongHeaderPacket && mrand.IntN(20) != 0 { + return false + } + numCorrupted.Add(1) + pos := mrand.IntN(len(b)) + b[pos] = byte(mrand.IntN(256)) + switch direction { + case quicproxy.DirectionIncoming: + clientTransport.WriteTo(b, serverTransport.Conn.LocalAddr()) + case quicproxy.DirectionOutgoing: + serverTransport.WriteTo(b, clientTransport.Conn.LocalAddr()) + } + return true + } + + runMITMTest(t, serverTransport, clientTransport, rtt, dropCallback) + t.Logf("corrupted %d packets", numCorrupted.Load()) + require.NotZero(t, int(numCorrupted.Load())) +} + +func runMITMTest(t *testing.T, serverTr, clientTr *quic.Transport, rtt time.Duration, dropCb quicproxy.DropCallback) { + ln, err := serverTr.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { return rtt / 2 }, + DropPacket: dropCb, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(time.Second)) + defer cancel() + conn, err := clientTr.Dial(ctx, proxy.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + str, err := conn.OpenStreamSync(ctx) + require.NoError(t, err) + clientErrChan := make(chan error, 1) + go func() { + _, err := str.Write(PRData) + clientErrChan <- err + str.Close() + }() + + serverStr, err := serverConn.AcceptStream(ctx) + require.NoError(t, err) + serverErrChan := make(chan error, 1) + go func() { + defer close(serverErrChan) + if _, err := io.Copy(serverStr, serverStr); err != nil { + serverErrChan <- err + return + } + serverStr.Close() + }() + require.NoError(t, <-serverErrChan) + + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRData, data) + + select { + case err := <-clientErrChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } + select { + case err := <-serverErrChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestMITMForgedVersionNegotiationPacket(t *testing.T) { + serverTransport, clientTransport := getTransportsForMITMTest(t) + rtt := scaleDuration(10 * time.Millisecond) + + const supportedVersion protocol.Version = 42 + + var once sync.Once + delayCb := func(dir quicproxy.Direction, _, _ net.Addr, raw []byte) time.Duration { + if dir != quicproxy.DirectionIncoming { + return rtt / 2 + } + once.Do(func() { + hdr, _, _, err := wire.ParsePacket(raw) + if err != nil { + panic("failed to parse packet: " + err.Error()) + } + // create fake version negotiation packet with a fake supported version + packet := wire.ComposeVersionNegotiation( + protocol.ArbitraryLenConnectionID(hdr.SrcConnectionID.Bytes()), + protocol.ArbitraryLenConnectionID(hdr.DestConnectionID.Bytes()), + []protocol.Version{supportedVersion}, + ) + if _, err := serverTransport.WriteTo(packet, clientTransport.Conn.LocalAddr()); err != nil { + panic("failed to write packet: " + err.Error()) + } + }) + return rtt / 2 + } + + err := runMITMTestSuccessful(t, serverTransport, clientTransport, delayCb) + var vnErr *quic.VersionNegotiationError + require.ErrorAs(t, err, &vnErr) + require.Contains(t, vnErr.Theirs, supportedVersion) // might contain greased versions +} + +// times out, because client doesn't accept subsequent real retry packets from server +// as it has already accepted a retry. +// TODO: determine behavior when server does not send Retry packets +func TestMITMForgedRetryPacket(t *testing.T) { + serverTransport, clientTransport := getTransportsForMITMTest(t) + serverTransport.VerifySourceAddress = func(net.Addr) bool { return true } + rtt := scaleDuration(10 * time.Millisecond) + + var once sync.Once + delayCb := func(dir quicproxy.Direction, _, _ net.Addr, raw []byte) time.Duration { + hdr, _, _, err := wire.ParsePacket(raw) + if err != nil { + panic("failed to parse packet: " + err.Error()) + } + if dir == quicproxy.DirectionIncoming && hdr.Type == protocol.PacketTypeInitial { + once.Do(func() { + fakeSrcConnID := protocol.ParseConnectionID([]byte{0x12, 0x12, 0x12, 0x12, 0x12, 0x12, 0x12, 0x12}) + retryPacket := testutils.ComposeRetryPacket(fakeSrcConnID, hdr.SrcConnectionID, hdr.DestConnectionID, []byte("token"), hdr.Version) + if _, err := serverTransport.WriteTo(retryPacket, clientTransport.Conn.LocalAddr()); err != nil { + panic("failed to write packet: " + err.Error()) + } + }) + } + return rtt / 2 + } + err := runMITMTestSuccessful(t, serverTransport, clientTransport, delayCb) + var nerr net.Error + require.ErrorAs(t, err, &nerr) + require.True(t, nerr.Timeout()) +} + +func TestMITMForgedInitialPacket(t *testing.T) { + serverTransport, clientTransport := getTransportsForMITMTest(t) + rtt := scaleDuration(10 * time.Millisecond) + + var once sync.Once + delayCb := func(dir quicproxy.Direction, _, _ net.Addr, raw []byte) time.Duration { + if dir == quicproxy.DirectionIncoming { + hdr, _, _, err := wire.ParsePacket(raw) + if err != nil { + panic("failed to parse packet: " + err.Error()) + } + if hdr.Type != protocol.PacketTypeInitial { + return 0 + } + once.Do(func() { + initialPacket := testutils.ComposeInitialPacket( + hdr.DestConnectionID, + hdr.SrcConnectionID, + hdr.DestConnectionID, + nil, + nil, + protocol.PerspectiveServer, + hdr.Version, + ) + if _, err := serverTransport.WriteTo(initialPacket, clientTransport.Conn.LocalAddr()); err != nil { + panic("failed to write packet: " + err.Error()) + } + }) + } + return rtt / 2 + } + err := runMITMTestSuccessful(t, serverTransport, clientTransport, delayCb) + var nerr net.Error + require.ErrorAs(t, err, &nerr) + require.True(t, nerr.Timeout()) +} + +func TestMITMForgedInitialPacketWithAck(t *testing.T) { + serverTransport, clientTransport := getTransportsForMITMTest(t) + rtt := scaleDuration(10 * time.Millisecond) + + var once sync.Once + delayCb := func(dir quicproxy.Direction, _, _ net.Addr, raw []byte) time.Duration { + if dir == quicproxy.DirectionIncoming { + hdr, _, _, err := wire.ParsePacket(raw) + if err != nil { + panic("failed to parse packet: " + err.Error()) + } + if hdr.Type != protocol.PacketTypeInitial { + return 0 + } + once.Do(func() { + // Fake Initial with ACK for packet 2 (unsent) + ack := &wire.AckFrame{AckRanges: []wire.AckRange{{Smallest: 2, Largest: 2}}} + initialPacket := testutils.ComposeInitialPacket( + hdr.DestConnectionID, + hdr.SrcConnectionID, + hdr.DestConnectionID, + nil, + []wire.Frame{ack}, + protocol.PerspectiveServer, + hdr.Version, + ) + if _, err := serverTransport.WriteTo(initialPacket, clientTransport.Conn.LocalAddr()); err != nil { + panic("failed to write packet: " + err.Error()) + } + }) + } + return rtt / 2 + } + + err := runMITMTestSuccessful(t, serverTransport, clientTransport, delayCb) + var transportErr *quic.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, quic.ProtocolViolation, transportErr.ErrorCode) + require.Contains(t, transportErr.ErrorMessage, "received ACK for an unsent packet") +} + +func runMITMTestSuccessful(t *testing.T, serverTransport, clientTransport *quic.Transport, delayCb quicproxy.DelayCallback) error { + t.Helper() + ln, err := serverTransport.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer ln.Close() + + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: delayCb, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + ctx, cancel := context.WithTimeout(context.Background(), scaleDuration(50*time.Millisecond)) + defer cancel() + _, err = clientTransport.Dial(ctx, proxy.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + require.Error(t, err) + return err +} diff --git a/third_party/quic-go/integrationtests/self/mtu_test.go b/third_party/quic-go/integrationtests/self/mtu_test.go new file mode 100644 index 0000000..78adcc7 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/mtu_test.go @@ -0,0 +1,197 @@ +package self_test + +import ( + "bytes" + "context" + "fmt" + "io" + "net" + "sync" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func TestInitialPacketSize(t *testing.T) { + server := newUDPConnLocalhost(t) + client := newUDPConnLocalhost(t) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + quic.Dial(ctx, client, server.LocalAddr(), getTLSClientConfig(), getQuicConfig(&quic.Config{ + InitialPacketSize: 1337, + })) + }() + + buf := make([]byte, 2000) + n, _, err := server.ReadFrom(buf) + require.NoError(t, err) + require.Equal(t, 1337, n) + + cancel() + <-done +} + +func TestPathMTUDiscovery(t *testing.T) { + rtt := scaleDuration(5 * time.Millisecond) + const mtu = 1400 + + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + InitialPacketSize: 1234, + DisablePathMTUDiscovery: true, + EnableDatagrams: true, + }), + ) + require.NoError(t, err) + defer ln.Close() + + serverErrChan := make(chan error, 1) + go func() { + conn, err := ln.Accept(context.Background()) + if err != nil { + serverErrChan <- err + return + } + str, err := conn.AcceptStream(context.Background()) + if err != nil { + serverErrChan <- err + return + } + defer str.Close() + if _, err := io.Copy(str, str); err != nil { + serverErrChan <- err + return + } + }() + + var mx sync.Mutex + var maxPacketSizeServer int + var clientPacketSizes []int + proxy := &quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { return rtt / 2 }, + DropPacket: func(dir quicproxy.Direction, _, _ net.Addr, packet []byte) bool { + if len(packet) > mtu { + return true + } + mx.Lock() + defer mx.Unlock() + switch dir { + case quicproxy.DirectionIncoming: + clientPacketSizes = append(clientPacketSizes, len(packet)) + case quicproxy.DirectionOutgoing: + if len(packet) > maxPacketSizeServer { + maxPacketSizeServer = len(packet) + } + } + return false + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + // Make sure to use v4-only socket here. + // We can't reliably set the DF bit on dual-stack sockets on older versions of macOS (before Sequoia). + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer tr.Close() + + var eventRecorder events.Recorder + conn, err := tr.Dial( + context.Background(), + proxy.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + InitialPacketSize: protocol.MinInitialPacketSize, + EnableDatagrams: true, + Tracer: newTracer(&eventRecorder), + }), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + err = conn.SendDatagram(make([]byte, 2000)) + require.Error(t, err) + var datagramErr *quic.DatagramTooLargeError + require.ErrorAs(t, err, &datagramErr) + initialMaxDatagramSize := datagramErr.MaxDatagramPayloadSize + + str, err := conn.OpenStream() + require.NoError(t, err) + + clientErrChan := make(chan error, 1) + go func() { + data, err := io.ReadAll(str) + if err != nil { + clientErrChan <- err + return + } + if !bytes.Equal(data, PRDataLong) { + clientErrChan <- fmt.Errorf("echoed data doesn't match: %x", data) + return + } + clientErrChan <- nil + }() + + _, err = str.Write(PRDataLong) + require.NoError(t, err) + str.Close() + + select { + case err := <-clientErrChan: + require.NoError(t, err) + case err := <-serverErrChan: + t.Fatalf("server error: %v", err) + case <-time.After(20 * time.Second): + t.Fatal("timeout") + } + + err = conn.SendDatagram(make([]byte, 2000)) + require.Error(t, err) + require.ErrorAs(t, err, &datagramErr) + finalMaxDatagramSize := datagramErr.MaxDatagramPayloadSize + + mx.Lock() + defer mx.Unlock() + require.NotEmpty(t, eventRecorder.Events(qlog.MTUUpdated{})) + + var mtus []int + for _, ev := range eventRecorder.Events(qlog.MTUUpdated{}) { + mtus = append(mtus, ev.(qlog.MTUUpdated).Value) + } + + maxPacketSizeClient := mtus[len(mtus)-1] + t.Logf("max client packet size: %d, MTU: %d", maxPacketSizeClient, mtu) + t.Logf("max datagram size: initial: %d, final: %d", initialMaxDatagramSize, finalMaxDatagramSize) + t.Logf("max server packet size: %d, MTU: %d", maxPacketSizeServer, mtu) + + require.GreaterOrEqual(t, maxPacketSizeClient, mtu-25) + const maxDiff = 40 // this includes the 21 bytes for the short header, 16 bytes for the encryption tag, and framing overhead + require.GreaterOrEqual(t, int(initialMaxDatagramSize), protocol.MinInitialPacketSize-maxDiff) + require.GreaterOrEqual(t, int(finalMaxDatagramSize), maxPacketSizeClient-maxDiff) + // MTU discovery was disabled on the server side + require.Equal(t, 1234, maxPacketSizeServer) + + var numPacketsLargerThanDiscoveredMTU int + for _, s := range clientPacketSizes { + if s > maxPacketSizeClient { + numPacketsLargerThanDiscoveredMTU++ + } + } + // The client shouldn't have sent any packets larger than the MTU it discovered, + // except for at most one MTU probe packet. + require.LessOrEqual(t, numPacketsLargerThanDiscoveredMTU, 1) +} diff --git a/third_party/quic-go/integrationtests/self/multiplex_test.go b/third_party/quic-go/integrationtests/self/multiplex_test.go new file mode 100644 index 0000000..0cfb00e --- /dev/null +++ b/third_party/quic-go/integrationtests/self/multiplex_test.go @@ -0,0 +1,346 @@ +package self_test + +import ( + "bytes" + "context" + "crypto/rand" + "errors" + "fmt" + "io" + mrand "math/rand/v2" + "net" + "runtime" + "testing" + "time" + + "github.com/apernet/quic-go" + + "github.com/stretchr/testify/require" +) + +func runMultiplexTestServer(t *testing.T, ln *quic.Listener) { + t.Helper() + for { + conn, err := ln.Accept(context.Background()) + if err != nil { + return + } + str, err := conn.OpenUniStream() + require.NoError(t, err) + go func() { + defer str.Close() + _, err = str.Write(PRData) + require.NoError(t, err) + }() + + t.Cleanup(func() { conn.CloseWithError(0, "") }) + } +} + +func dialAndReceiveData(tr *quic.Transport, addr net.Addr) error { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := tr.Dial(ctx, addr, getTLSClientConfig(), getQuicConfig(nil)) + if err != nil { + return fmt.Errorf("error dialing: %w", err) + } + str, err := conn.AcceptUniStream(ctx) + if err != nil { + return fmt.Errorf("error accepting stream: %w", err) + } + data, err := io.ReadAll(str) + if err != nil { + return fmt.Errorf("error reading data: %w", err) + } + if !bytes.Equal(data, PRData) { + return fmt.Errorf("data mismatch: got %q, expected %q", data, PRData) + } + return nil +} + +func TestMultiplexesConnectionsToSameServer(t *testing.T) { + server, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + go runMultiplexTestServer(t, server) + + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr) + defer tr.Close() + + errChan1 := make(chan error, 1) + go func() { errChan1 <- dialAndReceiveData(tr, server.Addr()) }() + errChan2 := make(chan error, 1) + go func() { errChan2 <- dialAndReceiveData(tr, server.Addr()) }() + + select { + case err := <-errChan1: + require.NoError(t, err, "error dialing server 1") + case <-time.After(5 * time.Second): + t.Error("timeout waiting for done1 to close") + } + select { + case err := <-errChan2: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Error("timeout waiting for done2 to close") + } +} + +func TestMultiplexingToDifferentServers(t *testing.T) { + server1, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server1.Close() + go runMultiplexTestServer(t, server1) + + server2, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server2.Close() + go runMultiplexTestServer(t, server2) + + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr) + defer tr.Close() + + errChan1 := make(chan error, 1) + go func() { errChan1 <- dialAndReceiveData(tr, server1.Addr()) }() + errChan2 := make(chan error, 1) + go func() { errChan2 <- dialAndReceiveData(tr, server2.Addr()) }() + + select { + case err := <-errChan1: + require.NoError(t, err, "error dialing server 1") + case <-time.After(5 * time.Second): + t.Error("timeout waiting for done1 to close") + } + select { + case err := <-errChan2: + require.NoError(t, err, "error dialing server 2") + case <-time.After(5 * time.Second): + t.Error("timeout waiting for done2 to close") + } +} + +func TestMultiplexingConnectToSelf(t *testing.T) { + tr := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr) + defer tr.Close() + + server, err := tr.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + go runMultiplexTestServer(t, server) + + errChan := make(chan error, 1) + go func() { errChan <- dialAndReceiveData(tr, server.Addr()) }() + + select { + case err := <-errChan: + require.NoError(t, err, "error dialing server") + case <-time.After(5 * time.Second): + t.Error("timeout waiting for connection to close") + } +} + +func TestMultiplexingServerAndClientOnSameConn(t *testing.T) { + if runtime.GOOS == "linux" { + t.Skip("This test requires setting of iptables rules on Linux, see https://stackoverflow.com/questions/23859164/linux-udp-socket-sendto-operation-not-permitted.") + } + + tr1 := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr1) + defer tr1.Close() + server1, err := tr1.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server1.Close() + + tr2 := &quic.Transport{Conn: newUDPConnLocalhost(t)} + addTracer(tr2) + defer tr2.Close() + server2, err := tr2.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server2.Close() + + done1 := make(chan struct{}) + go func() { + defer close(done1) + dialAndReceiveData(tr2, server1.Addr()) + }() + + done2 := make(chan struct{}) + go func() { + defer close(done2) + dialAndReceiveData(tr1, server2.Addr()) + }() + + select { + case <-done1: + case <-time.After(5 * time.Second): + t.Error("timeout waiting for done1 to close") + } + select { + case <-done2: + case <-time.After(time.Second): + t.Error("timeout waiting for done2 to close") + } +} + +func TestMultiplexingNonQUICPackets(t *testing.T) { + const numPackets = 100 + + tr1 := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer tr1.Close() + addTracer(tr1) + server, err := tr1.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer server.Close() + + tr2 := &quic.Transport{Conn: newUDPConnLocalhost(t)} + defer tr2.Close() + addTracer(tr2) + + type nonQUICPacket struct { + b []byte + addr net.Addr + err error + } + rcvdPackets := make(chan nonQUICPacket, numPackets) + receiveCtx := t.Context() + // start receiving non-QUIC packets + go func() { + for { + b := make([]byte, 1024) + n, addr, err := tr2.ReadNonQUICPacket(receiveCtx, b) + if errors.Is(err, context.Canceled) { + return + } + rcvdPackets <- nonQUICPacket{b: b[:n], addr: addr, err: err} + } + }() + + ctx2, cancel2 := context.WithTimeout(context.Background(), time.Second) + defer cancel2() + conn, err := tr2.Dial(ctx2, server.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverConn, err := server.Accept(ctx2) + require.NoError(t, err) + serverStr, err := serverConn.OpenUniStream() + require.NoError(t, err) + + // send a non-QUIC packet every 100µs + const packetLen = 128 + errChanNonQUIC := make(chan error, 1) + sendNonQUICPacket := make(chan struct{}, 1) + go func() { + var seed [32]byte + rand.Read(seed[:]) + random := mrand.NewChaCha8(seed) + defer close(errChanNonQUIC) + var sentPackets int + for range sendNonQUICPacket { + b := make([]byte, packetLen) + random.Read(b[1:]) // keep the first byte set to 0, so it's not classified as a QUIC packet + _, err := tr1.WriteTo(b, tr2.Conn.LocalAddr()) + // The first sendmsg call on a new UDP socket sometimes errors on Linux. + // It's not clear why this happens. + // See https://github.com/golang/go/issues/63322. + if err != nil && sentPackets == 0 && runtime.GOOS == "linux" && isPermissionError(err) { + _, err = tr1.WriteTo(b, tr2.Conn.LocalAddr()) + } + if err != nil { + errChanNonQUIC <- err + return + } + sentPackets++ + } + }() + + sendQUICPacket := make(chan struct{}, 1) + errChanQUIC := make(chan error, 1) + var dataSent []byte + go func() { + defer close(errChanQUIC) + defer serverStr.Close() + + var seed [32]byte + rand.Read(seed[:]) + random := mrand.NewChaCha8(seed) + for range sendQUICPacket { + b := make([]byte, 1024) + random.Read(b) + if _, err := serverStr.Write(b); err != nil { + errChanQUIC <- err + return + } + dataSent = append(dataSent, b...) + } + }() + + dataChan := make(chan []byte, 1) + readErr := make(chan error, 1) + go func() { + str, err := conn.AcceptUniStream(ctx2) + if err != nil { + readErr <- err + return + } + data, err := io.ReadAll(str) + if err != nil { + readErr <- err + return + } + dataChan <- data + }() + + ticker := time.NewTicker(scaleDuration(200 * time.Microsecond)) + defer ticker.Stop() + for range numPackets { + sendNonQUICPacket <- struct{}{} + sendQUICPacket <- struct{}{} + <-ticker.C + } + close(sendNonQUICPacket) + close(sendQUICPacket) + + select { + case err := <-errChanNonQUIC: + require.NoError(t, err, "error sending non-QUIC packets") + case <-time.After(time.Second): + t.Fatalf("timeout waiting for non-QUIC packets to be sent") + } + select { + case err := <-errChanQUIC: + require.NoError(t, err, "error sending QUIC packets") + case <-time.After(time.Second): + t.Fatalf("timeout waiting for QUIC packets to be sent") + } + select { + case err := <-readErr: + require.NoError(t, err, "error reading stream data") + case dataRcvd := <-dataChan: + require.Equal(t, dataSent, dataRcvd, "stream data mismatch") + case <-time.After(time.Second): + t.Fatalf("timeout waiting for stream data to be read") + } + + // make sure we don't overflow the capacity of the channel + require.LessOrEqual(t, numPackets, cap(rcvdPackets), "too many non-QUIC packets sent: %d > %d", numPackets, cap(rcvdPackets)) + + // now receive these packets + minExpected := numPackets * 4 / 5 + timeout := time.After(time.Second) + var counter int + for counter < minExpected { + select { + case p := <-rcvdPackets: + require.Equal(t, tr1.Conn.LocalAddr(), p.addr, "non-QUIC packet received from wrong address") + require.Equal(t, packetLen, len(p.b), "non-QUIC packet incorrect length") + require.NoError(t, p.err, "error receiving non-QUIC packet") + counter++ + case <-timeout: + t.Fatalf("didn't receive enough non-QUIC packets: %d < %d", counter, minExpected) + } + } +} diff --git a/third_party/quic-go/integrationtests/self/nat_rebinding_test.go b/third_party/quic-go/integrationtests/self/nat_rebinding_test.go new file mode 100644 index 0000000..a5fa58b --- /dev/null +++ b/third_party/quic-go/integrationtests/self/nat_rebinding_test.go @@ -0,0 +1,126 @@ +package self_test + +import ( + "context" + "fmt" + "io" + "net" + "os" + "sync" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/stretchr/testify/require" +) + +func TestNATRebinding(t *testing.T) { + tr, tracer := newPacketTracer() + tlsConf := getTLSConfig() + f, err := os.Create("keylog.txt") + require.NoError(t, err) + defer f.Close() + tlsConf.KeyLogWriter = f + server, err := quic.Listen( + newUDPConnLocalhost(t), + tlsConf, + getQuicConfig(&quic.Config{ + Tracer: func(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer server.Close() + + newPath := newUDPConnLocalhost(t) + clientUDPConn := newUDPConnLocalhost(t) + + oldPathRTT := scaleDuration(10 * time.Millisecond) + newPathRTT := scaleDuration(20 * time.Millisecond) + proxy := quicproxy.Proxy{ + ServerAddr: server.Addr().(*net.UDPAddr), + Conn: newUDPConnLocalhost(t), + } + var mx sync.Mutex + var switchedPath bool + var dataTransferred int + proxy.DelayPacket = func(dir quicproxy.Direction, _, _ net.Addr, b []byte) time.Duration { + mx.Lock() + defer mx.Unlock() + + if dir == quicproxy.DirectionOutgoing { + dataTransferred += len(b) + if dataTransferred > len(PRData)/3 { + if !switchedPath { + if err := proxy.SwitchConn(clientUDPConn.LocalAddr().(*net.UDPAddr), newPath); err != nil { + panic(fmt.Sprintf("failed to switch connection: %s", err)) + } + switchedPath = true + } + } + } + if switchedPath { + return newPathRTT + } + return oldPathRTT + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, clientUDPConn, proxy.LocalAddr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + go func() { + str, err := serverConn.OpenUniStream() + require.NoError(t, err) + go func() { + defer str.Close() + _, err = str.Write(PRData) + require.NoError(t, err) + }() + }() + + str, err := conn.AcceptUniStream(ctx) + require.NoError(t, err) + str.SetReadDeadline(time.Now().Add(5 * time.Second)) + data, err := io.ReadAll(str) + require.NoError(t, err) + require.Equal(t, PRData, data) + conn.CloseWithError(0, "") + + // check that a PATH_CHALLENGE was sent + var pathChallenge [8]byte + var foundPathChallenge bool + for _, p := range tr.getSentShortHeaderPackets() { + for _, f := range p.frames { + switch fr := f.Frame.(type) { + case *qlog.PathChallengeFrame: + pathChallenge = fr.Data + foundPathChallenge = true + } + } + } + require.True(t, foundPathChallenge) + + // check that a PATH_RESPONSE with the correct data was received + var foundPathResponse bool + for _, p := range tr.getRcvdShortHeaderPackets() { + for _, f := range p.frames { + switch fr := f.Frame.(type) { + case *qlog.PathResponseFrame: + require.Equal(t, pathChallenge, fr.Data) + foundPathResponse = true + } + } + } + require.True(t, foundPathResponse) +} diff --git a/third_party/quic-go/integrationtests/self/packetization_test.go b/third_party/quic-go/integrationtests/self/packetization_test.go new file mode 100644 index 0000000..b6eb765 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/packetization_test.go @@ -0,0 +1,273 @@ +package self_test + +import ( + "context" + "fmt" + "io" + "net" + "os" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestACKBundling(t *testing.T) { + const numMsg = 100 + + serverCounter, serverTracer := newPacketTracer() + server, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + DisablePathMTUDiscovery: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return serverTracer }, + }), + ) + require.NoError(t, err) + defer server.Close() + + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: server.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { + return 5 * time.Millisecond + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + clientCounter, clientTracer := newPacketTracer() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + proxy.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + DisablePathMTUDiscovery: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return clientTracer }, + }), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + + serverErrChan := make(chan error, 1) + go func() { + defer close(serverErrChan) + conn, err := server.Accept(context.Background()) + if err != nil { + serverErrChan <- fmt.Errorf("accept failed: %w", err) + return + } + str, err := conn.AcceptStream(context.Background()) + if err != nil { + serverErrChan <- fmt.Errorf("accept stream failed: %w", err) + return + } + b := make([]byte, 1) + // Echo every byte received from the client. + for { + if _, err := str.Read(b); err != nil { + break + } + _, err = str.Write(b) + if err != nil { + serverErrChan <- fmt.Errorf("write failed: %w", err) + return + } + } + }() + + str, err := conn.OpenStreamSync(context.Background()) + require.NoError(t, err) + b := make([]byte, 1) + // Send numMsg 1-byte messages. + for i := range numMsg { + _, err = str.Write([]byte{uint8(i)}) + require.NoError(t, err) + _, err = str.Read(b) + require.NoError(t, err) + require.Equal(t, uint8(i), b[0]) + } + require.NoError(t, conn.CloseWithError(0, "")) + require.NoError(t, <-serverErrChan) + + countBundledPackets := func(packets []packet) (numBundled int) { + for _, p := range packets { + var hasAck, hasStreamFrame bool + for _, f := range p.frames { + switch f.Frame.(type) { + case *qlog.AckFrame: + hasAck = true + case *qlog.StreamFrame: + hasStreamFrame = true + } + } + if hasAck && hasStreamFrame { + numBundled++ + } + } + return + } + + numBundledIncoming := countBundledPackets(clientCounter.getRcvdShortHeaderPackets()) + numBundledOutgoing := countBundledPackets(serverCounter.getRcvdShortHeaderPackets()) + t.Logf("bundled incoming packets: %d / %d", numBundledIncoming, numMsg) + t.Logf("bundled outgoing packets: %d / %d", numBundledOutgoing, numMsg) + + require.LessOrEqual(t, numBundledIncoming, numMsg) + require.Greater(t, numBundledIncoming, numMsg*9/10) + require.LessOrEqual(t, numBundledOutgoing, numMsg) + require.Greater(t, numBundledOutgoing, numMsg*9/10) +} + +func TestStreamDataBlocked(t *testing.T) { + testConnAndStreamDataBlocked(t, true, false) +} + +func TestConnDataBlocked(t *testing.T) { + testConnAndStreamDataBlocked(t, false, true) +} + +func testConnAndStreamDataBlocked(t *testing.T, limitStream, limitConn bool) { + const window = 100 + const numBatches = 3 + + initialStreamWindow := uint64(quicvarint.Max) + initialConnWindow := uint64(quicvarint.Max) + if limitStream { + initialStreamWindow = window + } + if limitConn { + initialConnWindow = window + } + rtt := scaleDuration(5 * time.Millisecond) + + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + InitialStreamReceiveWindow: initialStreamWindow, + InitialConnectionReceiveWindow: initialConnWindow, + }), + ) + require.NoError(t, err) + defer ln.Close() + + proxy := quicproxy.Proxy{ + Conn: newUDPConnLocalhost(t), + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { + return rtt / 2 + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + counter, tracer := newPacketTracer() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + proxy.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + + str, err := conn.OpenUniStreamSync(ctx) + require.NoError(t, err) + + // Stream data is consumed (almost) immediately, so flow-control window auto-tuning kicks in. + // The window size is doubled for every batch. + var windowSizes []protocol.ByteCount + for i := range numBatches { + windowSizes = append(windowSizes, window< highestSeen { + highestSeen = pn + } + } + + t.Logf("Smoothed RTT: %s", conn.ConnectionStats().SmoothedRTT) + assert.GreaterOrEqual(t, conn.ConnectionStats().SmoothedRTT, rtt*9/10) + assert.LessOrEqual(t, conn.ConnectionStats().SmoothedRTT, rtt*11/10) + t.Logf("received %d short header packets, detected %d reorderings", len(packetNumbers), reorderings) + assert.GreaterOrEqual(t, reorderings, 20, "expected at least 20 reorderings") + }) + }) + } +} diff --git a/third_party/quic-go/integrationtests/self/self_go124_test.go b/third_party/quic-go/integrationtests/self/self_go124_test.go new file mode 100644 index 0000000..16146e7 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/self_go124_test.go @@ -0,0 +1,9 @@ +//go:build !go1.25 + +package self_test + +import "crypto/tls" + +func getCurveID(connState tls.ConnectionState) tls.CurveID { + return 0 +} diff --git a/third_party/quic-go/integrationtests/self/self_go125_test.go b/third_party/quic-go/integrationtests/self/self_go125_test.go new file mode 100644 index 0000000..3474e15 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/self_go125_test.go @@ -0,0 +1,9 @@ +//go:build go1.25 + +package self_test + +import "crypto/tls" + +func getCurveID(connState tls.ConnectionState) tls.CurveID { + return connState.CurveID +} diff --git a/third_party/quic-go/integrationtests/self/self_suite_linux_test.go b/third_party/quic-go/integrationtests/self/self_suite_linux_test.go new file mode 100644 index 0000000..7baa1f4 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/self_suite_linux_test.go @@ -0,0 +1,21 @@ +//go:build linux + +package self_test + +import ( + "errors" + "os" + + "golang.org/x/sys/unix" +) + +// The first sendmsg call on a new UDP socket sometimes errors on Linux. +// It's not clear why this happens. +// See https://github.com/golang/go/issues/63322. +func isPermissionError(err error) bool { + var serr *os.SyscallError + if errors.As(err, &serr) { + return serr.Syscall == "sendmsg" && serr.Err == unix.EPERM + } + return false +} diff --git a/third_party/quic-go/integrationtests/self/self_suite_others_test.go b/third_party/quic-go/integrationtests/self/self_suite_others_test.go new file mode 100644 index 0000000..c0080d3 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/self_suite_others_test.go @@ -0,0 +1,7 @@ +//go:build !linux + +package self_test + +func isPermissionError(err error) bool { + return false +} diff --git a/third_party/quic-go/integrationtests/self/self_test.go b/third_party/quic-go/integrationtests/self/self_test.go new file mode 100644 index 0000000..e65f7c8 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/self_test.go @@ -0,0 +1,352 @@ +package self_test + +import ( + "context" + "crypto/tls" + "crypto/x509" + "flag" + "fmt" + "io" + "math/rand/v2" + "net" + "os" + "runtime" + "strconv" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/integrationtests/tools" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +const alpn = tools.ALPN + +const ( + dataLen = 500 * 1024 // 500 KB + dataLenLong = 50 * 1024 * 1024 // 50 MB +) + +var ( + // PRData contains dataLen bytes of pseudo-random data. + PRData = GeneratePRData(dataLen) + // PRDataLong contains dataLenLong bytes of pseudo-random data. + PRDataLong = GeneratePRData(dataLenLong) +) + +// See https://en.wikipedia.org/wiki/Lehmer_random_number_generator +func GeneratePRData(l int) []byte { + res := make([]byte, l) + seed := uint64(1) + for i := range l { + seed = seed * 48271 % 2147483647 + res[i] = byte(seed) + } + return res +} + +var ( + version quic.Version + enableQlog bool + + tlsConfig *tls.Config + tlsConfigLongChain *tls.Config + tlsClientConfig *tls.Config + tlsClientConfigWithoutServerName *tls.Config +) + +func init() { + ca, caPrivateKey, err := tools.GenerateCA() + if err != nil { + panic(err) + } + leafCert, leafPrivateKey, err := tools.GenerateLeafCert(ca, caPrivateKey) + if err != nil { + panic(err) + } + tlsConfig = &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: [][]byte{leafCert.Raw}, + PrivateKey: leafPrivateKey, + }}, + NextProtos: []string{alpn}, + } + tlsConfLongChain, err := tools.GenerateTLSConfigWithLongCertChain(ca, caPrivateKey) + if err != nil { + panic(err) + } + tlsConfigLongChain = tlsConfLongChain + + root := x509.NewCertPool() + root.AddCert(ca) + tlsClientConfig = &tls.Config{ + ServerName: "localhost", + RootCAs: root, + NextProtos: []string{alpn}, + } + tlsClientConfigWithoutServerName = &tls.Config{ + RootCAs: root, + NextProtos: []string{alpn}, + } +} + +func getTLSConfig() *tls.Config { return tlsConfig.Clone() } +func getTLSConfigWithLongCertChain() *tls.Config { return tlsConfigLongChain.Clone() } +func getTLSClientConfig() *tls.Config { return tlsClientConfig.Clone() } +func getTLSClientConfigWithoutServerName() *tls.Config { + return tlsClientConfigWithoutServerName.Clone() +} + +type multiplexedRecorder struct { + Recorders []qlogwriter.Recorder +} + +var _ qlogwriter.Recorder = &multiplexedRecorder{} + +func (r *multiplexedRecorder) Close() error { + for _, recorder := range r.Recorders { + recorder.Close() + } + return nil +} + +func (r *multiplexedRecorder) RecordEvent(ev qlogwriter.Event) { + for _, recorder := range r.Recorders { + recorder.RecordEvent(ev) + } +} + +type multiplexedTrace struct { + Traces []qlogwriter.Trace +} + +var _ qlogwriter.Trace = &multiplexedTrace{} + +func (t *multiplexedTrace) AddProducer() qlogwriter.Recorder { + recorders := make([]qlogwriter.Recorder, 0, len(t.Traces)) + for _, tr := range t.Traces { + recorders = append(recorders, tr.AddProducer()) + } + return &multiplexedRecorder{Recorders: recorders} +} + +func (t *multiplexedTrace) SupportsSchemas(schema string) bool { + return true +} + +func getQuicConfig(conf *quic.Config) *quic.Config { + if conf == nil { + conf = &quic.Config{} + } else { + conf = conf.Clone() + } + if !enableQlog { + return conf + } + if conf.Tracer == nil { + conf.Tracer = func(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + return tools.NewQlogConnectionTracer(os.Stdout)(ctx, isClient, connID) + } + return conf + } + origTracer := conf.Tracer + conf.Tracer = func(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + tr := origTracer(ctx, isClient, connID) + qlogger := tools.NewQlogConnectionTracer(os.Stdout)(ctx, isClient, connID) + if tr == nil { + return qlogger + } + return &multiplexedTrace{Traces: []qlogwriter.Trace{tr, qlogger}} + } + return conf +} + +func addTracer(tr *quic.Transport) { + if !enableQlog { + return + } + if tr.Tracer == nil { + tr.Tracer = tools.QlogTracer(os.Stdout).AddProducer() + return + } + origTracer := tr.Tracer + tr.Tracer = &multiplexedRecorder{ + Recorders: []qlogwriter.Recorder{origTracer, tools.QlogTracer(os.Stdout).AddProducer()}, + } +} + +func newUDPConnLocalhost(t testing.TB) *net.UDPConn { + t.Helper() + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + t.Cleanup(func() { conn.Close() }) + return conn +} + +func TestMain(m *testing.M) { + var versionParam string + flag.StringVar(&versionParam, "version", "1", "QUIC version") + flag.BoolVar(&enableQlog, "qlog", false, "enable qlog") + flag.Parse() + + switch versionParam { + case "1": + version = quic.Version1 + case "2": + version = quic.Version2 + default: + fmt.Printf("unknown QUIC version: %s\n", versionParam) + os.Exit(1) + } + fmt.Printf("using QUIC version: %s\n", version) + + os.Exit(m.Run()) +} + +func scaleDuration(d time.Duration) time.Duration { + scaleFactor := 1 + if f, err := strconv.Atoi(os.Getenv("TIMESCALE_FACTOR")); err == nil { // parsing "" errors, so this works fine if the env is not set + scaleFactor = f + } + if scaleFactor == 0 { + panic("TIMESCALE_FACTOR is 0") + } + return time.Duration(scaleFactor) * d +} + +func newTracer(tracer qlogwriter.Recorder) func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: tracer} + } +} + +type packet struct { + time time.Time + hdr qlog.PacketHeader + frames []qlog.Frame +} + +type packetCounter struct { + recorder *events.Recorder +} + +func (t *packetCounter) getSentShortHeaderPackets() []packet { + var sentShortHdr []packet + for _, ev := range t.recorder.EventsWithTime(qlog.PacketSent{}) { + e := ev.Event.(qlog.PacketSent) + if e.Header.PacketType != qlog.PacketType1RTT { + continue + } + sentShortHdr = append(sentShortHdr, packet{time: ev.Time, hdr: e.Header, frames: e.Frames}) + } + return sentShortHdr +} + +func (t *packetCounter) getRcvdLongHeaderPackets() []packet { + var rcvdLongHdr []packet + for _, ev := range t.recorder.EventsWithTime(qlog.PacketReceived{}) { + e := ev.Event.(qlog.PacketReceived) + if e.Header.PacketType == qlog.PacketType1RTT { + continue + } + rcvdLongHdr = append(rcvdLongHdr, packet{time: ev.Time, hdr: e.Header, frames: e.Frames}) + } + return rcvdLongHdr +} + +func (t *packetCounter) getRcvd0RTTPacketNumbers() []protocol.PacketNumber { + var zeroRTTPackets []protocol.PacketNumber + for _, p := range t.getRcvdLongHeaderPackets() { + if p.hdr.PacketType == qlog.PacketType0RTT { + zeroRTTPackets = append(zeroRTTPackets, p.hdr.PacketNumber) + } + } + return zeroRTTPackets +} + +func (t *packetCounter) getRcvdShortHeaderPackets() []packet { + var rcvdShortHdr []packet + for _, ev := range t.recorder.EventsWithTime(qlog.PacketReceived{}) { + e := ev.Event.(qlog.PacketReceived) + if e.Header.PacketType != qlog.PacketType1RTT { + continue + } + rcvdShortHdr = append(rcvdShortHdr, packet{time: ev.Time, hdr: e.Header, frames: e.Frames}) + } + return rcvdShortHdr +} + +func newPacketTracer() (*packetCounter, qlogwriter.Trace) { + c := &packetCounter{recorder: &events.Recorder{}} + return c, &events.Trace{Recorder: c.recorder} +} + +type readerWithTimeout struct { + io.Reader + Timeout time.Duration +} + +func (r *readerWithTimeout) Read(p []byte) (n int, err error) { + done := make(chan struct{}) + go func() { + defer close(done) + n, err = r.Reader.Read(p) + }() + + select { + case <-done: + return n, err + case <-time.After(r.Timeout): + return 0, fmt.Errorf("read timeout after %s", r.Timeout) + } +} + +func randomDuration(min, max time.Duration) time.Duration { + return min + time.Duration(rand.IntN(int(max-min))) +} + +// containsPacketType checks if a packet contains a long header packet of the specified type. +// It correctly handles coalesced packets. +func containsPacketType(data []byte, packetType protocol.PacketType) bool { + for len(data) > 0 { + if !wire.IsLongHeaderPacket(data[0]) { + return false + } + hdr, _, rest, err := wire.ParsePacket(data) + if err != nil { + return false + } + if hdr.Type == packetType { + return true + } + data = rest + } + return false +} + +// addDialCallback explicitly adds the http3.Transport's Dial callback. +// This is needed since dialing on dual-stack sockets is flaky on macOS, +// see https://github.com/golang/go/issues/67226. +func addDialCallback(t *testing.T, tr *http3.Transport) { + t.Helper() + + if runtime.GOOS != "darwin" { + return + } + + require.Nil(t, tr.Dial) + tr.Dial = func(ctx context.Context, addr string, tlsConf *tls.Config, conf *quic.Config) (*quic.Conn, error) { + a, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return nil, err + } + return quic.DialEarly(ctx, newUDPConnLocalhost(t), a, tlsConf, conf) + } +} diff --git a/third_party/quic-go/integrationtests/self/simnet_helper_test.go b/third_party/quic-go/integrationtests/self/simnet_helper_test.go new file mode 100644 index 0000000..ed150ee --- /dev/null +++ b/third_party/quic-go/integrationtests/self/simnet_helper_test.go @@ -0,0 +1,105 @@ +package self_test + +import ( + "net" + "testing" + "time" + + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/require" +) + +func newSimnetLink(t *testing.T, rtt time.Duration) (client, server *simnet.SimConn, close func(t *testing.T)) { + t.Helper() + + return newSimnetLinkWithRouter(t, rtt, &simnet.PerfectRouter{}) +} + +func newSimnetLinkWithRouter(t *testing.T, rtt time.Duration, router simnet.Router) (client, server *simnet.SimConn, close func(t *testing.T)) { + t.Helper() + + n := &simnet.Simnet{Router: router} + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + clientPacketConn := n.NewEndpoint(&net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001}, settings) + serverPacketConn := n.NewEndpoint(&net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002}, settings) + + require.NoError(t, n.Start()) + + return clientPacketConn, serverPacketConn, func(t *testing.T) { + require.NoError(t, clientPacketConn.Close()) + require.NoError(t, serverPacketConn.Close()) + require.NoError(t, n.Close()) + } +} + +type droppingRouter struct { + simnet.PerfectRouter + + Drop func(simnet.Packet) bool +} + +func (d *droppingRouter) SendPacket(p simnet.Packet) error { + if d.Drop(p) { + return nil + } + return d.PerfectRouter.SendPacket(p) +} + +type callbackRouter struct { + simnet.Router + + OnSendPacket func(simnet.Packet) +} + +func (c *callbackRouter) SendPacket(p simnet.Packet) error { + c.OnSendPacket(p) + return c.Router.SendPacket(p) +} + +type direction uint8 + +const ( + directionUnknown = iota + directionToClient + directionToServer + directionBoth +) + +func (d direction) String() string { + switch d { + case directionToClient: + return "to client" + case directionToServer: + return "to server" + case directionBoth: + return "both" + } + return "unknown" +} + +var _ simnet.Router = &droppingRouter{} + +type directionAwareDroppingRouter struct { + simnet.PerfectRouter + + ClientAddr, ServerAddr *net.UDPAddr + + Drop func(direction direction, p simnet.Packet) bool +} + +func (d *directionAwareDroppingRouter) SendPacket(p simnet.Packet) error { + var dir direction + switch p.To.String() { + case d.ClientAddr.String(): + dir = directionToClient + case d.ServerAddr.String(): + dir = directionToServer + default: + dir = directionUnknown + } + if d.Drop(dir, p) { + return nil + } + return d.PerfectRouter.SendPacket(p) +} diff --git a/third_party/quic-go/integrationtests/self/stateless_reset_test.go b/third_party/quic-go/integrationtests/self/stateless_reset_test.go new file mode 100644 index 0000000..4542247 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/stateless_reset_test.go @@ -0,0 +1,127 @@ +package self_test + +import ( + "context" + "crypto/rand" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/require" +) + +func TestStatelessResets(t *testing.T) { + t.Run("zero-length connection IDs", func(t *testing.T) { + testStatelessReset(t, 0) + }) + t.Run("10 byte connection IDs", func(t *testing.T) { + testStatelessReset(t, 10) + }) +} + +func testStatelessReset(t *testing.T, connIDLen int) { + synctest.Test(t, func(t *testing.T) { + var drop atomic.Bool + clientPacketConn, serverPacketConn, closeFn := newSimnetLinkWithRouter(t, + time.Millisecond, + &droppingRouter{Drop: func(p simnet.Packet) bool { return drop.Load() }}, + ) + defer closeFn(t) + + var statelessResetKey quic.StatelessResetKey + rand.Read(statelessResetKey[:]) + + tr := &quic.Transport{ + Conn: serverPacketConn, + StatelessResetKey: &statelessResetKey, + } + defer tr.Close() + + ln, err := tr.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + serverErr := make(chan error, 1) + go func() { + conn, err := ln.Accept(context.Background()) + if err != nil { + serverErr <- err + return + } + str, err := conn.OpenStream() + if err != nil { + serverErr <- err + return + } + _, err = str.Write([]byte("foobar")) + if err != nil { + serverErr <- err + return + } + close(serverErr) + }() + + var conn *quic.Conn + if connIDLen > 0 { + cl := &quic.Transport{ + Conn: clientPacketConn, + ConnectionIDLength: connIDLen, + } + defer cl.Close() + var err error + conn, err = cl.Dial( + context.Background(), + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIdleTimeout: 2 * time.Second}), + ) + require.NoError(t, err) + } else { + conn, err = quic.Dial( + context.Background(), + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIdleTimeout: 2 * time.Second}), + ) + require.NoError(t, err) + } + str, err := conn.AcceptStream(context.Background()) + require.NoError(t, err) + data := make([]byte, 6) + _, err = str.Read(data) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), data) + + // make sure that the CONNECTION_CLOSE is dropped + drop.Store(true) + require.NoError(t, ln.Close()) + require.NoError(t, tr.Close()) + require.NoError(t, <-serverErr) + time.Sleep(100 * time.Millisecond) + + // We need to create a new Transport here, since the old one is still sending out + // CONNECTION_CLOSE packets for (recently) closed connections). + tr2 := &quic.Transport{ + Conn: serverPacketConn, + StatelessResetKey: &statelessResetKey, + } + defer tr2.Close() + ln2, err := tr2.Listen(getTLSConfig(), getQuicConfig(nil)) + require.NoError(t, err) + drop.Store(false) + + // Trigger something (not too small) to be sent, so that we receive the stateless reset. + // If the client already sent another packet, it might already have received a packet. + _, serr := str.Write([]byte("Lorem ipsum dolor sit amet.")) + if serr == nil { + _, serr = str.Read([]byte{0}) + } + require.Error(t, serr) + require.IsType(t, &quic.StatelessResetError{}, serr) + require.NoError(t, ln2.Close()) + }) +} diff --git a/third_party/quic-go/integrationtests/self/stream_test.go b/third_party/quic-go/integrationtests/self/stream_test.go new file mode 100644 index 0000000..91a07f9 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/stream_test.go @@ -0,0 +1,375 @@ +package self_test + +import ( + "bytes" + "context" + "fmt" + "io" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/quicvarint" + + "golang.org/x/sync/errgroup" + + "github.com/stretchr/testify/require" +) + +func TestWriteWithLimitStreamFlowControl(t *testing.T) { + testWriteWithLimitFlowControl(t, &quic.Config{ + InitialStreamReceiveWindow: 100, + InitialConnectionReceiveWindow: quicvarint.Max, + }) +} + +func TestWriteWithLimitConnectionFlowControl(t *testing.T) { + testWriteWithLimitFlowControl(t, &quic.Config{ + InitialStreamReceiveWindow: quicvarint.Max, + InitialConnectionReceiveWindow: 100, + }) +} + +func testWriteWithLimitFlowControl(t *testing.T, config *quic.Config) { + ln, err := quic.Listen(newUDPConnLocalhost(t), getTLSConfig(), getQuicConfig(config)) + require.NoError(t, err) + defer ln.Close() + + deadline := time.Now().Add(time.Second) + ctx, cancel := context.WithDeadline(context.Background(), deadline) + defer cancel() + client, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + defer client.CloseWithError(0, "") + + server, err := ln.Accept(ctx) + require.NoError(t, err) + str, err := client.OpenUniStreamSync(ctx) + require.NoError(t, err) + require.NoError(t, str.SetWriteDeadline(deadline)) + + data := GeneratePRData(101) + n, err := str.WriteWithLimit(data, func(maxBytes int) int { + return min(maxBytes, 25) + }) + require.Equal(t, 25, n) + require.ErrorIs(t, err, quic.ErrWriteLimitReached) + + receiveStr, err := server.AcceptUniStream(ctx) + require.NoError(t, err) + require.NoError(t, receiveStr.SetReadDeadline(deadline)) + received := make([]byte, n) + _, err = io.ReadFull(receiveStr, received) + require.NoError(t, err) + require.Equal(t, data[:n], received) + + written, err := str.WriteWithLimit(data[n:], func(maxBytes int) int { return maxBytes }) + require.Equal(t, len(data)-n, written) + require.NoError(t, err) + + require.NoError(t, str.Close()) + rest, err := io.ReadAll(receiveStr) + require.NoError(t, err) + require.Equal(t, data, append(received, rest...)) +} + +func TestBidirectionalStreamMultiplexing(t *testing.T) { + const numStreams = 75 + + runSendingPeer := func(conn *quic.Conn) error { + g := new(errgroup.Group) + for i := range numStreams { + str, err := conn.OpenStreamSync(context.Background()) + if err != nil { + return err + } + data := GeneratePRData(50 * i) + g.Go(func() error { + if _, err := str.Write(data); err != nil { + return err + } + return str.Close() + }) + g.Go(func() error { + dataRead, err := io.ReadAll(str) + if err != nil { + return err + } + if !bytes.Equal(dataRead, data) { + return fmt.Errorf("data mismatch: %q != %q", dataRead, data) + } + return nil + }) + } + return g.Wait() + } + + runReceivingPeer := func(conn *quic.Conn) error { + g := new(errgroup.Group) + for range numStreams { + str, err := conn.AcceptStream(context.Background()) + if err != nil { + return err + } + g.Go(func() error { + // shouldn't use io.Copy here + // we should read from the stream as early as possible, to free flow control credit + data, err := io.ReadAll(str) + if err != nil { + return err + } + if _, err := str.Write(data); err != nil { + return err + } + return str.Close() + }) + } + return g.Wait() + } + + t.Run("client -> server", func(t *testing.T) { + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + MaxIncomingStreams: 10, + InitialStreamReceiveWindow: 10000, + InitialConnectionReceiveWindow: 5000, + }), + ) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + client, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + ln.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{InitialConnectionReceiveWindow: 2000}), + ) + require.NoError(t, err) + + conn, err := ln.Accept(ctx) + require.NoError(t, err) + + errChan := make(chan error, 1) + go func() { errChan <- runReceivingPeer(conn) }() + require.NoError(t, runSendingPeer(client)) + client.CloseWithError(0, "") + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + select { + case <-conn.Context().Done(): + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + }) + + t.Run("client <-> server", func(t *testing.T) { + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(&quic.Config{ + MaxIncomingStreams: 30, + InitialStreamReceiveWindow: 25000, + InitialConnectionReceiveWindow: 50000, + }), + ) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + client, err := quic.Dial( + ctx, + newUDPConnLocalhost(t), + ln.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{InitialConnectionReceiveWindow: 2000}), + ) + require.NoError(t, err) + + conn, err := ln.Accept(ctx) + require.NoError(t, err) + + errChan1 := make(chan error, 1) + errChan2 := make(chan error, 1) + errChan3 := make(chan error, 1) + errChan4 := make(chan error, 1) + + go func() { errChan1 <- runReceivingPeer(conn) }() + go func() { errChan2 <- runSendingPeer(conn) }() + go func() { errChan3 <- runReceivingPeer(client) }() + go func() { errChan4 <- runSendingPeer(client) }() + + for _, ch := range []chan error{errChan1, errChan2, errChan3, errChan4} { + select { + case err := <-ch: + require.NoError(t, err) + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + } + + client.CloseWithError(0, "") + select { + case <-conn.Context().Done(): + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + }) +} + +func TestUnidirectionalStreams(t *testing.T) { + const numStreams = 500 + + dataForStream := func(id uint64) []byte { return GeneratePRData(10 * int(id)) } + + runSendingPeer := func(conn *quic.Conn) error { + g := new(errgroup.Group) + for range numStreams { + str, err := conn.OpenUniStreamSync(context.Background()) + if err != nil { + return err + } + g.Go(func() error { + if _, err := str.Write(dataForStream(uint64(str.StreamID()))); err != nil { + return err + } + return str.Close() + }) + } + return g.Wait() + } + + runReceivingPeer := func(conn *quic.Conn) error { + g := new(errgroup.Group) + for range numStreams { + str, err := conn.AcceptUniStream(context.Background()) + if err != nil { + return err + } + g.Go(func() error { + data, err := io.ReadAll(str) + if err != nil { + return err + } + if !bytes.Equal(data, dataForStream(uint64(str.StreamID()))) { + return fmt.Errorf("data mismatch") + } + return nil + }) + } + return g.Wait() + } + + t.Run("client -> server", func(t *testing.T) { + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(nil), + ) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + client, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { errChan <- runSendingPeer(client) }() + require.NoError(t, runReceivingPeer(serverConn)) + serverConn.CloseWithError(0, "") + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + }) + + t.Run("server -> client", func(t *testing.T) { + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(nil), + ) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + client, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { errChan <- runSendingPeer(serverConn) }() + + require.NoError(t, runReceivingPeer(client)) + client.CloseWithError(0, "") + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + }) + + t.Run("client <-> server", func(t *testing.T) { + ln, err := quic.Listen( + newUDPConnLocalhost(t), + getTLSConfig(), + getQuicConfig(nil), + ) + require.NoError(t, err) + defer ln.Close() + + errChan1 := make(chan error, 1) + errChan2 := make(chan error, 1) + go func() { + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan1 <- err + errChan2 <- err + return + } + errChan1 <- runReceivingPeer(conn) + errChan2 <- runSendingPeer(conn) + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + client, err := quic.Dial(ctx, newUDPConnLocalhost(t), ln.Addr(), getTLSClientConfig(), getQuicConfig(nil)) + require.NoError(t, err) + + errChan3 := make(chan error, 1) + go func() { + errChan3 <- runSendingPeer(client) + }() + require.NoError(t, runReceivingPeer(client)) + + for _, ch := range []chan error{errChan1, errChan2, errChan3} { + select { + case err := <-ch: + require.NoError(t, err) + case <-time.After(time.Second): + require.Fail(t, "timeout") + } + } + client.CloseWithError(0, "") + }) +} diff --git a/third_party/quic-go/integrationtests/self/timeout_test.go b/third_party/quic-go/integrationtests/self/timeout_test.go new file mode 100644 index 0000000..f42d1f4 --- /dev/null +++ b/third_party/quic-go/integrationtests/self/timeout_test.go @@ -0,0 +1,504 @@ +package self_test + +import ( + "bytes" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + mrand "math/rand/v2" + "net" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/require" +) + +func requireIdleTimeoutError(t *testing.T, err error) { + t.Helper() + + require.Error(t, err) + var idleTimeoutErr *quic.IdleTimeoutError + require.ErrorAs(t, err, &idleTimeoutErr) + require.True(t, idleTimeoutErr.Timeout()) + var nerr net.Error + require.True(t, errors.As(err, &nerr)) + require.True(t, nerr.Timeout()) +} + +func TestHandshakeIdleTimeout(t *testing.T) { + t.Run("Dial", func(t *testing.T) { + testHandshakeIdleTimeout(t, quic.Dial) + }) + + t.Run("DialEarly", func(t *testing.T) { + testHandshakeIdleTimeout(t, quic.DialEarly) + }) +} + +func testHandshakeIdleTimeout(t *testing.T, dialFn func(context.Context, net.PacketConn, net.Addr, *tls.Config, *quic.Config) (*quic.Conn, error)) { + synctest.Test(t, func(t *testing.T) { + const handshakeIdleTimeout = 3 * time.Second + + clientPacketConn, serverPacketConn, closeFn := newSimnetLink(t, time.Millisecond) + defer closeFn(t) + + errChan := make(chan error, 1) + start := time.Now() + go func() { + _, err := dialFn( + context.Background(), + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{HandshakeIdleTimeout: handshakeIdleTimeout}), + ) + errChan <- err + }() + select { + case err := <-errChan: + requireIdleTimeoutError(t, err) + require.Equal(t, handshakeIdleTimeout, time.Since(start)) + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for dial error") + } + }) +} + +func TestIdleTimeout(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const idleTimeout = 20 * time.Second + + var drop atomic.Bool + clientPacketConn, serverPacketConn, closeFn := newSimnetLinkWithRouter(t, + time.Millisecond, + &droppingRouter{Drop: func(p simnet.Packet) bool { return drop.Load() }}, + ) + defer closeFn(t) + + server, err := quic.Listen( + serverPacketConn, + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + defer server.Close() + + conn, err := quic.Dial( + context.Background(), + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true, MaxIdleTimeout: idleTimeout}), + ) + require.NoError(t, err) + + serverConn, err := server.Accept(context.Background()) + require.NoError(t, err) + str, err := serverConn.OpenStream() + require.NoError(t, err) + _, err = str.Write([]byte("foobar")) + require.NoError(t, err) + + serverStart := time.Now() + + strIn, err := conn.AcceptStream(context.Background()) + require.NoError(t, err) + strOut, err := conn.OpenStream() + require.NoError(t, err) + _, err = strIn.Read(make([]byte, 6)) + require.NoError(t, err) + + clientStart := time.Now() + + drop.Store(true) + + select { + case <-serverConn.Context().Done(): + took := time.Since(serverStart) + require.GreaterOrEqual(t, took, idleTimeout) + t.Logf("server connection timed out after %s (idle timeout: %s)", took, idleTimeout) + case <-time.After(2 * idleTimeout): + t.Fatal("timeout waiting for idle timeout") + } + + select { + case <-conn.Context().Done(): + took := time.Since(clientStart) + require.GreaterOrEqual(t, took, idleTimeout) + t.Logf("client connection timed out after %s (idle timeout: %s)", took, idleTimeout) + case <-time.After(2 * idleTimeout): + t.Fatal("timeout waiting for idle timeout") + } + + _, err = strIn.Write([]byte("test")) + requireIdleTimeoutError(t, err) + _, err = strIn.Read([]byte{0}) + requireIdleTimeoutError(t, err) + _, err = strOut.Write([]byte("test")) + requireIdleTimeoutError(t, err) + _, err = strOut.Read([]byte{0}) + requireIdleTimeoutError(t, err) + _, err = conn.OpenStream() + requireIdleTimeoutError(t, err) + _, err = conn.OpenUniStream() + requireIdleTimeoutError(t, err) + _, err = conn.AcceptStream(context.Background()) + requireIdleTimeoutError(t, err) + _, err = conn.AcceptUniStream(context.Background()) + requireIdleTimeoutError(t, err) + }) +} + +func TestKeepAlive(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const idleTimeout = 4 * time.Second + + var drop atomic.Bool + clientPacketConn, serverPacketConn, closeFn := newSimnetLinkWithRouter(t, + time.Millisecond, + &droppingRouter{Drop: func(p simnet.Packet) bool { return drop.Load() }}, + ) + defer closeFn(t) + + server, err := quic.Listen( + serverPacketConn, + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial( + ctx, + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + MaxIdleTimeout: idleTimeout, + KeepAlivePeriod: idleTimeout / 2, + DisablePathMTUDiscovery: true, + }), + ) + require.NoError(t, err) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + + // wait longer than the idle timeout + time.Sleep(3 * idleTimeout) + str, err := conn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write([]byte("foobar")) + require.NoError(t, err) + + // verify connection is still alive + select { + case <-serverConn.Context().Done(): + t.Fatal("server connection closed unexpectedly") + default: + } + + // idle timeout will still kick in if PINGs are dropped + drop.Store(true) + time.Sleep(2 * idleTimeout) + _, err = str.Write([]byte("foobar")) + requireIdleTimeoutError(t, err) + + // can't rely on the server connection closing, since we impose a minimum idle timeout of 5s, + // see https://github.com/apernet/quic-go/issues/4751 + serverConn.CloseWithError(0, "") + }) +} + +func TestTimeoutAfterInactivity(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const idleTimeout = 15 * time.Second + + clientPacketConn, serverPacketConn, closeFn := newSimnetLink(t, time.Millisecond) + defer closeFn(t) + + server, err := quic.Listen( + serverPacketConn, + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + counter, tr := newPacketTracer() + conn, err := quic.Dial( + ctx, + clientPacketConn, + server.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{ + MaxIdleTimeout: idleTimeout, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tr }, + DisablePathMTUDiscovery: true, + }), + ) + require.NoError(t, err) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + defer serverConn.CloseWithError(0, "") + + ctx, cancel = context.WithTimeout(context.Background(), 2*idleTimeout) + defer cancel() + _, err = conn.AcceptStream(ctx) + requireIdleTimeoutError(t, err) + + var lastAckElicitingPacketSentAt time.Time + for _, p := range counter.getSentShortHeaderPackets() { + var hasAckElicitingFrame bool + for _, f := range p.frames { + if _, ok := f.Frame.(qlog.AckFrame); ok { + continue + } + hasAckElicitingFrame = true + break + } + if hasAckElicitingFrame { + lastAckElicitingPacketSentAt = p.time + } + } + rcvdPackets := counter.getRcvdShortHeaderPackets() + lastPacketRcvdAt := rcvdPackets[len(rcvdPackets)-1].time + // We're ignoring here that only the first ack-eliciting packet sent resets the idle timeout. + // This is ok since we're dealing with a lossless connection here, + // and we'd expect to receive an ACK for additional other ack-eliciting packet sent. + timeSinceLastAckEliciting := time.Since(lastAckElicitingPacketSentAt) + timeSinceLastRcvd := time.Since(lastPacketRcvdAt) + require.Equal(t, idleTimeout, max(timeSinceLastAckEliciting, timeSinceLastRcvd)) + + select { + case <-serverConn.Context().Done(): + t.Fatal("server connection closed unexpectedly") + default: + } + }) +} + +func TestTimeoutAfterSendingPacket(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const idleTimeout = 15 * time.Second + + var drop atomic.Bool + clientPacketConn, serverPacketConn, closeFn := newSimnetLinkWithRouter(t, + time.Millisecond, + &droppingRouter{Drop: func(p simnet.Packet) bool { return drop.Load() }}, + ) + defer closeFn(t) + + server, err := quic.Listen( + serverPacketConn, + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + defer server.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial( + ctx, + clientPacketConn, + serverPacketConn.LocalAddr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{MaxIdleTimeout: idleTimeout, DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + + serverConn, err := server.Accept(ctx) + require.NoError(t, err) + + serverStart := time.Now() + + // wait half the idle timeout, then send a packet + time.Sleep(idleTimeout / 2) + drop.Store(true) + + clientStart := time.Now() + str, err := conn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write([]byte("foobar")) + require.NoError(t, err) + + select { + case <-serverConn.Context().Done(): + took := time.Since(serverStart) + require.GreaterOrEqual(t, took, idleTimeout) + require.Less(t, took, idleTimeout+time.Second) + case <-time.After(2 * idleTimeout): + t.Fatal("timeout waiting for idle timeout") + } + + select { + case <-conn.Context().Done(): + took := time.Since(clientStart) + require.Equal(t, took, idleTimeout) + case <-time.After(2 * idleTimeout): + t.Fatal("timeout waiting for idle timeout") + } + }) +} + +type faultyConn struct { + net.PacketConn + + MaxPackets int + counter atomic.Int32 +} + +func (c *faultyConn) ReadFrom(p []byte) (int, net.Addr, error) { + n, addr, err := c.PacketConn.ReadFrom(p) + counter := c.counter.Add(1) + if counter <= int32(c.MaxPackets) { + return n, addr, err + } + return 0, nil, io.ErrClosedPipe +} + +func (c *faultyConn) WriteTo(p []byte, addr net.Addr) (int, error) { + counter := c.counter.Add(1) + if counter <= int32(c.MaxPackets) { + return c.PacketConn.WriteTo(p, addr) + } + return 0, io.ErrClosedPipe +} + +func TestFaultyPacketConn(t *testing.T) { + t.Run("client", func(t *testing.T) { + testFaultyPacketConn(t, protocol.PerspectiveClient) + }) + + t.Run("server", func(t *testing.T) { + testFaultyPacketConn(t, protocol.PerspectiveServer) + }) +} + +func testFaultyPacketConn(t *testing.T, pers protocol.Perspective) { + t.Setenv("QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING", "true") + + synctest.Test(t, func(t *testing.T) { + runServer := func(ln *quic.Listener) error { + conn, err := ln.Accept(context.Background()) + if err != nil { + return err + } + str, err := conn.OpenUniStream() + if err != nil { + return err + } + defer str.Close() + _, err = str.Write(PRData) + return err + } + + runClient := func(conn *quic.Conn) error { + str, err := conn.AcceptUniStream(context.Background()) + if err != nil { + return err + } + data, err := io.ReadAll(str) + if err != nil { + return err + } + if !bytes.Equal(data, PRData) { + return fmt.Errorf("wrong data: %q vs %q", data, PRData) + } + return conn.CloseWithError(0, "done") + } + + clientPacketConn, serverPacketConn, closeFn := newSimnetLink(t, 100*time.Millisecond) + defer closeFn(t) + + var cconn, sconn net.PacketConn = clientPacketConn, serverPacketConn + maxPackets := mrand.IntN(25) + // sanity check: sending PRData should generate at least 25 packets + require.Greater(t, len(PRData)/1500, 25) + + t.Logf("blocking %s's connection after %d packets", pers, maxPackets) + switch pers { + case protocol.PerspectiveClient: + cconn = &faultyConn{PacketConn: cconn, MaxPackets: maxPackets} + case protocol.PerspectiveServer: + sconn = &faultyConn{PacketConn: sconn, MaxPackets: maxPackets} + } + + ln, err := quic.Listen( + sconn, + getTLSConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + require.NoError(t, err) + defer ln.Close() + + serverErrChan := make(chan error, 1) + go func() { serverErrChan <- runServer(ln) }() + + clientErrChan := make(chan error, 1) + go func() { + conn, err := quic.Dial( + context.Background(), + cconn, + ln.Addr(), + getTLSClientConfig(), + getQuicConfig(&quic.Config{DisablePathMTUDiscovery: true}), + ) + if err != nil { + clientErrChan <- err + return + } + clientErrChan <- runClient(conn) + }() + + var clientErr error + select { + case clientErr = <-clientErrChan: + case <-time.After(time.Hour): + t.Fatal("timeout waiting for client error") + } + require.Error(t, clientErr) + if pers == protocol.PerspectiveClient { + require.Contains(t, clientErr.Error(), io.ErrClosedPipe.Error()) + } else { + var nerr net.Error + require.True(t, errors.As(clientErr, &nerr)) + require.True(t, nerr.Timeout()) + } + + select { + case serverErr := <-serverErrChan: // The handshake completed on the server side. + require.Error(t, serverErr) + if pers == protocol.PerspectiveServer { + require.Contains(t, serverErr.Error(), io.ErrClosedPipe.Error()) + } else { + var nerr net.Error + require.True(t, errors.As(serverErr, &nerr)) + require.True(t, nerr.Timeout()) + } + default: // The handshake didn't complete + require.NoError(t, ln.Close()) + select { + case <-serverErrChan: + case <-time.After(time.Hour): + t.Fatal("timeout waiting for server to close") + } + } + }) +} diff --git a/third_party/quic-go/integrationtests/self/zero_rtt_test.go b/third_party/quic-go/integrationtests/self/zero_rtt_test.go new file mode 100644 index 0000000..2a4fddf --- /dev/null +++ b/third_party/quic-go/integrationtests/self/zero_rtt_test.go @@ -0,0 +1,1154 @@ +package self_test + +import ( + "bytes" + "context" + "crypto/tls" + "fmt" + "io" + "net" + "os" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/require" +) + +type zeroRTTCountingRouter struct { + simnet.Router + counter atomic.Uint32 +} + +var _ simnet.Router = &zeroRTTCountingRouter{} + +func (r *zeroRTTCountingRouter) SendPacket(p simnet.Packet) error { + if containsPacketType(p.Data, protocol.PacketType0RTT) { + r.counter.Add(1) + } + return r.Router.SendPacket(p) +} + +func (r *zeroRTTCountingRouter) Num0RTTPackets() int { + return int(r.counter.Load()) +} + +func dialAndReceiveTicket(t *testing.T, ln *quic.EarlyListener, clientConn net.PacketConn, sessionCache tls.ClientSessionCache) (clientTLSConf *tls.Config) { + t.Helper() + + clientTLSConf = getTLSClientConfig() + puts := make(chan string, 100) + cache := sessionCache + if cache == nil { + cache = tls.NewLRUClientSessionCache(100) + } + clientTLSConf.ClientSessionCache = newClientSessionCache(cache, nil, puts) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + tr := &quic.Transport{Conn: clientConn} + defer tr.Close() + conn, err := tr.Dial(ctx, ln.Addr(), clientTLSConf, getQuicConfig(nil)) + require.NoError(t, err) + require.False(t, conn.ConnectionState().Used0RTT) + + select { + case <-puts: + case <-time.After(time.Second): + t.Fatal("timeout waiting for session ticket") + } + require.NoError(t, conn.CloseWithError(0, "")) + + serverConn, err := ln.Accept(ctx) + require.NoError(t, err) + + select { + case <-serverConn.Context().Done(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for connection to close") + } + return clientTLSConf +} + +func transfer0RTTData( + t *testing.T, + ln *quic.EarlyListener, + clientPacketConn net.PacketConn, + clientTLSConf *tls.Config, + clientConf *quic.Config, + testdata []byte, // data to transfer +) { + t.Helper() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + tr := &quic.Transport{Conn: clientPacketConn} + defer tr.Close() + conn, err := tr.DialEarly(ctx, ln.Addr(), clientTLSConf, clientConf) + require.NoError(t, err) + + errChan := make(chan error, 1) + serverConnChan := make(chan *quic.Conn, 1) + go func() { + defer close(errChan) + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan <- err + return + } + serverConnChan <- conn + str, err := conn.AcceptStream(ctx) + if err != nil { + errChan <- err + return + } + defer str.Close() + if _, err := io.Copy(str, str); err != nil { + errChan <- err + return + } + }() + + str, err := conn.OpenStream() + require.NoError(t, err) + + clientErrChan := make(chan error, 1) + go func() { + defer close(clientErrChan) + // wait for the EOF from the server to arrive before closing the conn + data, err := io.ReadAll(str) + if err != nil { + t.Error(err) + clientErrChan <- err + return + } + if !bytes.Equal(testdata, data) { + clientErrChan <- fmt.Errorf("data mismatch") + } + }() + + _, err = str.Write(testdata) + require.NoError(t, err) + require.NoError(t, str.Close()) + select { + case <-conn.HandshakeComplete(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for handshake to complete") + } + + select { + case err := <-clientErrChan: + require.NoError(t, err) + case <-time.After(time.Hour): + t.Fatal("timeout waiting for client to read data") + } + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout waiting for server to process data") + } + + var serverConn *quic.Conn + select { + case serverConn = <-serverConnChan: + case <-time.After(time.Second): + t.Fatal("timeout waiting for server to process data") + } + + require.True(t, conn.ConnectionState().Used0RTT) + require.True(t, serverConn.ConnectionState().Used0RTT) + conn.CloseWithError(0, "") + + select { + case <-serverConn.Context().Done(): + case <-time.After(time.Second): + t.Fatal("timeout waiting for connection to close") + } +} + +func Test0RTTTransfer(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 50 * time.Millisecond + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + counter, tracer := newPacketTracer() + defer tr.Close() + ln, err := tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + + time.Sleep(time.Hour) + synctest.Wait() + + transfer0RTTData(t, ln, clientConn, clientTLSConf, getQuicConfig(nil), PRData) + + num0RTT := router.Num0RTTPackets() + t.Logf("sent %d 0-RTT packets", num0RTT) + zeroRTTPackets := counter.getRcvd0RTTPacketNumbers() + t.Logf("received %d 0-RTT packets", len(zeroRTTPackets)) + require.Greater(t, num0RTT, 20) + require.Contains(t, zeroRTTPackets, protocol.PacketNumber(0)) + }) +} + +func Test0RTTDisabledOnDial(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 25 * time.Millisecond + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true})) + require.NoError(t, err) + defer ln.Close() + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + + time.Sleep(time.Hour) + synctest.Wait() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.Dial(ctx, clientConn, serverConn.LocalAddr(), clientTLSConf, getQuicConfig(nil)) + require.NoError(t, err) + // session Resumption is enabled at the TLS layer, but not 0-RTT at the QUIC layer + require.True(t, conn.ConnectionState().TLS.DidResume) + require.False(t, conn.ConnectionState().Used0RTT) + conn.CloseWithError(0, "") + + require.Zero(t, router.Num0RTTPackets()) + }) +} + +func Test0RTTWaitForHandshakeCompletion(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 50 * time.Millisecond + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + counter, tracer := newPacketTracer() + ln, err := tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + + zeroRTTData := GeneratePRData(5 << 10) + oneRTTData := PRData + + // now accept the second connection, and receive the 0-RTT data + errChan := make(chan error, 1) + firstStrDataChan := make(chan []byte, 1) + secondStrDataChan := make(chan []byte, 1) + go func() { + defer close(errChan) + conn, err := ln.Accept(context.Background()) + if err != nil { + errChan <- err + return + } + str, err := conn.AcceptUniStream(context.Background()) + if err != nil { + errChan <- err + return + } + data, err := io.ReadAll(str) + if err != nil { + errChan <- err + return + } + firstStrDataChan <- data + str, err = conn.AcceptUniStream(context.Background()) + if err != nil { + errChan <- err + return + } + data, err = io.ReadAll(str) + if err != nil { + errChan <- err + return + } + secondStrDataChan <- data + <-conn.Context().Done() + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.DialEarly( + ctx, + clientConn, + serverConn.LocalAddr(), + clientTLSConf, + getQuicConfig(nil), + ) + require.NoError(t, err) + firstStr, err := conn.OpenUniStream() + require.NoError(t, err) + _, err = firstStr.Write(zeroRTTData) + require.NoError(t, err) + require.NoError(t, firstStr.Close()) + + // wait for the handshake to complete + select { + case <-conn.HandshakeComplete(): + case <-time.After(time.Second): + t.Fatal("handshake did not complete in time") + } + str, err := conn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write(PRData) + require.NoError(t, err) + require.NoError(t, str.Close()) + + select { + case data := <-firstStrDataChan: + require.Equal(t, zeroRTTData, data) + case <-time.After(time.Second): + t.Fatal("timeout waiting for first stream data") + } + select { + case data := <-secondStrDataChan: + require.Equal(t, oneRTTData, data) + case <-time.After(time.Second): + t.Fatal("timeout waiting for second stream data") + } + conn.CloseWithError(0, "") + select { + case err := <-errChan: + require.NoError(t, err, "server error") + case <-time.After(time.Second): + t.Fatal("timeout waiting for connection to close") + } + + // check that 0-RTT packets only contain STREAM frames for the first stream + var num0RTT int + for _, p := range counter.getRcvdLongHeaderPackets() { + if p.hdr.PacketType != qlog.PacketType0RTT { + continue + } + for _, f := range p.frames { + sf, ok := f.Frame.(*qlog.StreamFrame) + if !ok { + continue + } + num0RTT++ + require.Equal(t, firstStr.StreamID(), sf.StreamID) + } + } + t.Logf("received %d STREAM frames in 0-RTT packets", num0RTT) + require.NotZero(t, num0RTT) + }) +} + +func Test0RTTDataLoss(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + tlsConf := getTLSConfig() + + var num0RTTPackets, numDropped atomic.Uint32 + router := &droppingRouter{ + Drop: func(p simnet.Packet) bool { + if !wire.IsLongHeaderPacket(p.Data[0]) { + return false + } + hdr, _, _, _ := wire.ParsePacket(p.Data) + if hdr.Type == protocol.PacketType0RTT { + count := num0RTTPackets.Add(1) + // drop 25% of the 0-RTT packets + drop := count%4 == 0 + if drop { + numDropped.Add(1) + } + return drop + } + return false + }, + } + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + counter, tracer := newPacketTracer() + ln, err := tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + + transfer0RTTData(t, ln, clientConn, clientTLSConf, getQuicConfig(nil), PRData) + + num0RTT := num0RTTPackets.Load() + dropped := numDropped.Load() + t.Logf("sent %d 0-RTT packets, dropped %d of those.", num0RTT, dropped) + require.NotZero(t, num0RTT) + require.NotZero(t, dropped) + require.NotEmpty(t, counter.getRcvd0RTTPacketNumbers()) + }) +} + +func Test0RTTRetransmitOnRetry(t *testing.T) { + t.Run("no retry", func(t *testing.T) { + test0RTTRetransmitOnRetry(t, false) + }) + t.Run("with retry", func(t *testing.T) { + test0RTTRetransmitOnRetry(t, true) + }) +} + +func test0RTTRetransmitOnRetry(t *testing.T, useRetry bool) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + tlsConf := getTLSConfig() + + type connIDCounter struct { + connID protocol.ConnectionID + bytes protocol.ByteCount + } + var mutex sync.Mutex + var connIDToCounter []*connIDCounter + countZeroRTTBytes := func(data []byte) (n protocol.ByteCount) { + for len(data) > 0 { + hdr, _, rest, err := wire.ParsePacket(data) + if err != nil { + return + } + data = rest + if hdr.Type == protocol.PacketType0RTT { + n += hdr.Length - 16 /* AEAD tag */ + } + } + return + } + + router := &zeroRTTCountingRouter{ + Router: &callbackRouter{ + Router: &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}}, + OnSendPacket: func(p simnet.Packet) { + if l := countZeroRTTBytes(p.Data); l > 0 { + mutex.Lock() + defer mutex.Unlock() + + connID, err := wire.ParseConnectionID(p.Data, 0) + if err != nil { + panic("failed to parse connection ID") + } + var found bool + for _, c := range connIDToCounter { + if c.connID == connID { + c.bytes += l + found = true + break + } + } + if !found { + connIDToCounter = append(connIDToCounter, &connIDCounter{connID: connID, bytes: l}) + } + } + }, + }, + } + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tr := &quic.Transport{ + Conn: serverConn, + VerifySourceAddress: func(net.Addr) bool { return useRetry }, + } + defer tr.Close() + counter, tracer := newPacketTracer() + ln, err := tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.Empty(t, connIDToCounter) + + transfer0RTTData(t, ln, clientConn, clientTLSConf, getQuicConfig(nil), GeneratePRData(5000)) // ~5 packets + + mutex.Lock() + defer mutex.Unlock() + + if !useRetry { + require.Len(t, connIDToCounter, 1) + return + } + + require.Len(t, connIDToCounter, 2) + require.InDelta(t, 5000+100 /* framing overhead */, int(connIDToCounter[0].bytes), 100) // the FIN bit might be sent extra + require.InDelta(t, int(connIDToCounter[0].bytes), int(connIDToCounter[1].bytes), 20) + zeroRTTPackets := counter.getRcvd0RTTPacketNumbers() + require.GreaterOrEqual(t, len(zeroRTTPackets), 5) + require.GreaterOrEqual(t, zeroRTTPackets[0], protocol.PacketNumber(5)) + }) +} + +func Test0RTTWithIncreasedStreamLimit(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true, MaxIncomingUniStreams: 1})) + require.NoError(t, err) + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.Zero(t, router.Num0RTTPackets()) + require.NoError(t, ln.Close()) + + time.Sleep(time.Hour) + synctest.Wait() + + ln, err = tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true, MaxIncomingUniStreams: 2})) + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.DialEarly( + ctx, + clientConn, + ln.Addr(), + clientTLSConf, + getQuicConfig(nil), + ) + require.NoError(t, err) + require.False(t, conn.ConnectionState().TLS.HandshakeComplete) + str, err := conn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, str.Close()) + // the client remembers the old limit and refuses to open a new stream + _, err = conn.OpenUniStream() + require.ErrorIs(t, err, &quic.StreamLimitReachedError{}) + + // after handshake completion, the new limit applies + select { + case <-conn.HandshakeComplete(): + case <-time.After(time.Second): + t.Fatal("handshake did not complete in time") + } + _, err = conn.OpenUniStream() + require.NoError(t, err) + require.True(t, conn.ConnectionState().Used0RTT) + require.NoError(t, conn.CloseWithError(0, "")) + + require.NotZero(t, router.Num0RTTPackets()) + }) +} + +func check0RTTRejected(t *testing.T, + ln *quic.EarlyListener, + clientPacketConn net.PacketConn, + addr net.Addr, + conf *tls.Config, + sendData bool, +) (clientConn, serverConn *quic.Conn) { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.DialEarly(ctx, clientPacketConn, addr, conf, getQuicConfig(nil)) + require.NoError(t, err) + require.False(t, conn.ConnectionState().TLS.HandshakeComplete) + if sendData { + str, err := conn.OpenUniStream() + require.NoError(t, err) + _, err = str.Write(make([]byte, 3000)) + require.NoError(t, err) + require.NoError(t, str.Close()) + } + + select { + case <-conn.HandshakeComplete(): + case <-time.After(time.Second): + t.Fatal("handshake did not complete in time") + } + require.False(t, conn.ConnectionState().Used0RTT) + + // make sure the server doesn't process the data + ctx, cancel = context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverConn, err = ln.Accept(ctx) + require.NoError(t, err) + require.False(t, serverConn.ConnectionState().Used0RTT) + if sendData { + _, err = serverConn.AcceptUniStream(ctx) + require.Equal(t, context.DeadlineExceeded, err) + } + + ctx, cancel = context.WithTimeout(context.Background(), time.Second) + defer cancel() + nextConn, err := conn.NextConnection(ctx) + require.NoError(t, err) + require.True(t, nextConn.ConnectionState().TLS.HandshakeComplete) + require.False(t, nextConn.ConnectionState().Used0RTT) + return nextConn, serverConn +} + +func Test0RTTRejectedOnStreamLimitDecrease(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + + const ( + maxBidiStreams = 42 + maxUniStreams = 10 + newMaxBidiStreams = maxBidiStreams - 1 + newMaxUniStreams = maxUniStreams - 1 + ) + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + MaxIncomingStreams: maxBidiStreams, + MaxIncomingUniStreams: maxUniStreams, + }), + ) + require.NoError(t, err) + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.NoError(t, ln.Close()) + + time.Sleep(time.Hour) + synctest.Wait() + + counter, tracer := newPacketTracer() + ln, err = tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + MaxIncomingStreams: newMaxBidiStreams, + MaxIncomingUniStreams: newMaxUniStreams, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + + conn, sconn := check0RTTRejected(t, ln, clientConn, ln.Addr(), clientTLSConf, true) + defer conn.CloseWithError(0, "") + + // It should now be possible to open new bidirectional streams up to the new limit... + for range newMaxBidiStreams { + _, err = conn.OpenStream() + require.NoError(t, err) + } + // ... but not beyond it. + _, err = conn.OpenStream() + require.ErrorIs(t, err, &quic.StreamLimitReachedError{}) + + // It should now be possible to open new unidirectional streams up to the new limit... + for range newMaxUniStreams { + _, err = conn.OpenUniStream() + require.NoError(t, err) + } + // ... but not beyond it. + _, err = conn.OpenUniStream() + require.ErrorIs(t, err, &quic.StreamLimitReachedError{}) + + sconn.CloseWithError(0, "") + // The client should send 0-RTT packets, but the server doesn't process them. + n := router.Num0RTTPackets() + t.Logf("sent %d 0-RTT packets", n) + require.NotZero(t, n) + require.Empty(t, counter.getRcvd0RTTPacketNumbers()) + }) +} + +func Test0RTTRejectedOnConnectionWindowDecrease(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + + const ( + connFlowControlWindow = 100 + newConnFlowControlWindow = connFlowControlWindow - 1 + ) + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + InitialConnectionReceiveWindow: connFlowControlWindow, + }), + ) + require.NoError(t, err) + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.NoError(t, ln.Close()) + + time.Sleep(time.Hour) + synctest.Wait() + + ln, err = tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + InitialConnectionReceiveWindow: newConnFlowControlWindow, + }), + ) + require.NoError(t, err) + + conn, sconn := check0RTTRejected(t, ln, clientConn, ln.Addr(), clientTLSConf, false) + defer conn.CloseWithError(0, "") + defer sconn.CloseWithError(0, "") + + str, err := conn.OpenStream() + require.NoError(t, err) + str.SetWriteDeadline(time.Now().Add(time.Second)) + n, err := str.Write(make([]byte, 2000)) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Equal(t, newConnFlowControlWindow, n) + + // make sure that only 99 bytes were received + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + serverStr, err := sconn.AcceptStream(ctx) + require.NoError(t, err) + serverStr.SetReadDeadline(time.Now().Add(time.Second)) + n, err = io.ReadFull(serverStr, make([]byte, newConnFlowControlWindow)) + require.NoError(t, err) + require.Equal(t, newConnFlowControlWindow, n) + _, err = serverStr.Read([]byte{0}) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + }) +} + +func Test0RTTRejectedOnALPNChanged(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true})) + require.NoError(t, err) + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.NoError(t, ln.Close()) + + time.Sleep(time.Hour) + synctest.Wait() + + // switch to different ALPN on the server side + tlsConf.NextProtos = []string{"new-alpn"} + // Append to the client's ALPN. + // crypto/tls will attempt to resume with the ALPN from the original connection + clientTLSConf.NextProtos = append(clientTLSConf.NextProtos, "new-alpn") + counter, tracer := newPacketTracer() + ln, err = tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + + conn, sconn := check0RTTRejected(t, ln, clientConn, ln.Addr(), clientTLSConf, true) + defer conn.CloseWithError(0, "") + + require.Equal(t, "new-alpn", conn.ConnectionState().TLS.NegotiatedProtocol) + + sconn.CloseWithError(0, "") + // The client should send 0-RTT packets, but the server doesn't process them. + num0RTT := router.Num0RTTPackets() + t.Logf("Sent %d 0-RTT packets.", num0RTT) + require.NotZero(t, num0RTT) + require.Empty(t, counter.getRcvd0RTTPacketNumbers()) + }) +} + +func Test0RTTRejectedWhenDisabled(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true})) + require.NoError(t, err) + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.NoError(t, ln.Close()) + + time.Sleep(time.Hour) + synctest.Wait() + + counter, tracer := newPacketTracer() + ln, err = tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: false, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + conn, sconn := check0RTTRejected(t, ln, clientConn, ln.Addr(), clientTLSConf, true) + defer conn.CloseWithError(0, "") + + sconn.CloseWithError(0, "") + // The client should send 0-RTT packets, but the server doesn't process them. + num0RTT := router.Num0RTTPackets() + t.Logf("Sent %d 0-RTT packets.", num0RTT) + require.NotZero(t, num0RTT) + require.Empty(t, counter.getRcvd0RTTPacketNumbers()) + }) +} + +func Test0RTTRejectedOnDatagramsDisabled(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tlsConf := getTLSConfig() + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true, EnableDatagrams: true})) + require.NoError(t, err) + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + require.NoError(t, ln.Close()) + + time.Sleep(time.Hour) + synctest.Wait() + + counter, tracer := newPacketTracer() + ln, err = tr.ListenEarly( + tlsConf, + getQuicConfig(&quic.Config{ + Allow0RTT: true, + EnableDatagrams: false, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + conn, sconn := check0RTTRejected(t, ln, clientConn, ln.Addr(), clientTLSConf, true) + defer conn.CloseWithError(0, "") + require.False(t, conn.ConnectionState().SupportsDatagrams.Remote) + require.False(t, conn.ConnectionState().SupportsDatagrams.Local) + + sconn.CloseWithError(0, "") + // The client should send 0-RTT packets, but the server doesn't process them. + num0RTT := router.Num0RTTPackets() + t.Logf("Sent %d 0-RTT packets.", num0RTT) + require.NotZero(t, num0RTT) + require.Empty(t, counter.getRcvd0RTTPacketNumbers()) + }) +} + +type metadataClientSessionCache struct { + toAdd []byte + restored func([]byte) + + cache tls.ClientSessionCache +} + +func (m metadataClientSessionCache) Get(key string) (*tls.ClientSessionState, bool) { + session, ok := m.cache.Get(key) + if !ok || session == nil { + return session, ok + } + ticket, state, err := session.ResumptionState() + if err != nil { + panic("failed to get resumption state: " + err.Error()) + } + if len(state.Extra) != 2 { // ours, and the quic-go's + panic("expected 2 state entries" + fmt.Sprintf("%v", state.Extra)) + } + m.restored(state.Extra[1]) + // as of Go 1.23, this function never returns an error + session, err = tls.NewResumptionState(ticket, state) + if err != nil { + panic("failed to create resumption state: " + err.Error()) + } + return session, true +} + +func (m metadataClientSessionCache) Put(key string, session *tls.ClientSessionState) { + ticket, state, err := session.ResumptionState() + if err != nil { + panic("failed to get resumption state: " + err.Error()) + } + state.Extra = append(state.Extra, m.toAdd) + session, err = tls.NewResumptionState(ticket, state) + if err != nil { + panic("failed to create resumption state: " + err.Error()) + } + m.cache.Put(key, session) +} + +func Test0RTTWithSessionTicketData(t *testing.T) { + const rtt = 5 * time.Millisecond + + t.Run("server", func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + tlsConf := getTLSConfig() + tlsConf.WrapSession = func(cs tls.ConnectionState, ss *tls.SessionState) ([]byte, error) { + ss.Extra = append(ss.Extra, []byte("foobar")) + return tlsConf.EncryptTicket(cs, ss) + } + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(tlsConf, getQuicConfig(&quic.Config{Allow0RTT: true})) + require.NoError(t, err) + + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + stateChan := make(chan *tls.SessionState, 1) + tlsConf.UnwrapSession = func(identity []byte, cs tls.ConnectionState) (*tls.SessionState, error) { + state, err := tlsConf.DecryptTicket(identity, cs) + if err != nil { + panic("failed to decrypt ticket") + } + stateChan <- state + return state, nil + } + + transfer0RTTData(t, ln, clientConn, clientTLSConf, getQuicConfig(nil), PRData) + + select { + case state := <-stateChan: + require.Len(t, state.Extra, 2) + require.Equal(t, []byte("foobar"), state.Extra[1]) + case <-time.After(time.Second): + t.Fatal("timed out waiting for session state") + } + }) + }) + + t.Run("client", func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + ln, err := tr.ListenEarly(getTLSConfig(), getQuicConfig(&quic.Config{Allow0RTT: true})) + require.NoError(t, err) + defer ln.Close() + + restoreChan := make(chan []byte, 1) + clientTLSConf := dialAndReceiveTicket(t, + ln, + clientConn, + &metadataClientSessionCache{ + toAdd: []byte("foobar"), + restored: func(b []byte) { restoreChan <- b }, + cache: tls.NewLRUClientSessionCache(100), + }, + ) + + transfer0RTTData(t, ln, clientConn, clientTLSConf, getQuicConfig(nil), PRData) + select { + case b := <-restoreChan: + require.Equal(t, []byte("foobar"), b) + case <-time.After(time.Second): + t.Fatal("timed out waiting for session state") + } + }) + }) +} + +func Test0RTTPacketQueueing(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + n := &simnet.Simnet{Router: &simnet.PerfectRouter{}} + serverAddr := &net.UDPAddr{IP: net.ParseIP("1.0.0.2"), Port: 9002} + settings := simnet.NodeBiDiLinkSettings{ + LatencyFunc: func(p simnet.Packet) time.Duration { + if p.To.String() == serverAddr.String() { + if wire.IsLongHeaderPacket(p.Data[0]) { + hdr, _, _, err := wire.ParsePacket(p.Data) + if err == nil && hdr.Type == protocol.PacketTypeInitial { + return rtt * 3 / 2 + } + } + } + return rtt / 2 + }, + } + clientConn := n.NewEndpoint(&net.UDPAddr{IP: net.ParseIP("1.0.0.1"), Port: 9001}, settings) + serverConn := n.NewEndpoint(serverAddr, settings) + require.NoError(t, n.Start()) + defer func() { + require.NoError(t, clientConn.Close()) + require.NoError(t, serverConn.Close()) + require.NoError(t, n.Close()) + }() + + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + counter, tracer := newPacketTracer() + ln, err := tr.ListenEarly( + getTLSConfig(), + getQuicConfig(&quic.Config{ + Allow0RTT: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + + data := GeneratePRData(5000) // ~5 packets + transfer0RTTData(t, ln, clientConn, clientTLSConf, getQuicConfig(nil), data) + + require.Equal(t, qlog.PacketTypeInitial, counter.getRcvdLongHeaderPackets()[0].hdr.PacketType) + zeroRTTPackets := counter.getRcvd0RTTPacketNumbers() + require.GreaterOrEqual(t, len(zeroRTTPackets), 5) + // make sure the data wasn't retransmitted + var dataSent protocol.ByteCount + for _, p := range counter.getRcvdLongHeaderPackets() { + for _, f := range p.frames { + if sf, ok := f.Frame.(*qlog.StreamFrame); ok { + dataSent += protocol.ByteCount(sf.Length) + } + } + } + for _, p := range counter.getRcvdShortHeaderPackets() { + for _, f := range p.frames { + if sf, ok := f.Frame.(*qlog.StreamFrame); ok { + dataSent += protocol.ByteCount(sf.Length) + } + } + } + require.Less(t, int(dataSent), 6000) + require.Equal(t, protocol.PacketNumber(0), zeroRTTPackets[0]) + }) +} + +func Test0RTTDatagrams(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 5 * time.Millisecond + router := &zeroRTTCountingRouter{Router: &simnet.PerfectRouter{}} + clientConn, serverConn, closeFn := newSimnetLinkWithRouter(t, rtt, router) + defer closeFn(t) + + tr := &quic.Transport{Conn: serverConn} + defer tr.Close() + + counter, tracer := newPacketTracer() + ln, err := tr.ListenEarly( + getTLSConfig(), + getQuicConfig(&quic.Config{ + Allow0RTT: true, + EnableDatagrams: true, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { return tracer }, + }), + ) + require.NoError(t, err) + defer ln.Close() + + clientTLSConf := dialAndReceiveTicket(t, ln, clientConn, nil) + + msg := GeneratePRData(100) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := quic.DialEarly(ctx, + clientConn, + ln.Addr(), + clientTLSConf, + getQuicConfig(&quic.Config{EnableDatagrams: true}), + ) + require.NoError(t, err) + defer conn.CloseWithError(0, "") + require.True(t, conn.ConnectionState().SupportsDatagrams.Remote) + require.True(t, conn.ConnectionState().SupportsDatagrams.Local) + require.NoError(t, conn.SendDatagram(msg)) + select { + case <-conn.HandshakeComplete(): + case <-time.After(time.Second): + t.Fatal("handshake did not complete in time") + } + + sconn, err := ln.Accept(ctx) + require.NoError(t, err) + rcvdMsg, err := sconn.ReceiveDatagram(ctx) + require.NoError(t, err) + require.True(t, sconn.ConnectionState().Used0RTT) + require.Equal(t, msg, rcvdMsg) + + num0RTT := router.Num0RTTPackets() + t.Logf("sent %d 0-RTT packets", num0RTT) + require.NotZero(t, num0RTT) + sconn.CloseWithError(0, "") + require.Len(t, counter.getRcvd0RTTPacketNumbers(), 1) + }) +} diff --git a/third_party/quic-go/integrationtests/tools/crypto.go b/third_party/quic-go/integrationtests/tools/crypto.go new file mode 100644 index 0000000..9c7bd37 --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/crypto.go @@ -0,0 +1,127 @@ +package tools + +import ( + "crypto" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "math/big" + "net" + "time" +) + +const ALPN = "quic-go integration tests" + +// use a very long validity period to cover the synthetic clock used in synctest +var ( + notBefore = time.Date(1990, 1, 1, 0, 0, 0, 0, time.UTC) + notAfter = time.Date(2100, 1, 1, 0, 0, 0, 0, time.UTC) +) + +func GenerateCA() (*x509.Certificate, crypto.PrivateKey, error) { + certTempl := &x509.Certificate{ + SerialNumber: big.NewInt(2019), + Subject: pkix.Name{}, + NotBefore: notBefore, + NotAfter: notAfter, + IsCA: true, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + BasicConstraintsValid: true, + } + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + caBytes, err := x509.CreateCertificate(rand.Reader, certTempl, certTempl, pub, priv) + if err != nil { + return nil, nil, err + } + ca, err := x509.ParseCertificate(caBytes) + if err != nil { + return nil, nil, err + } + return ca, priv, nil +} + +func GenerateLeafCert(ca *x509.Certificate, caPriv crypto.PrivateKey) (*x509.Certificate, crypto.PrivateKey, error) { + certTempl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + DNSNames: []string{"localhost"}, + IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)}, + NotBefore: notBefore, + NotAfter: notAfter, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, + KeyUsage: x509.KeyUsageDigitalSignature, + } + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + certBytes, err := x509.CreateCertificate(rand.Reader, certTempl, ca, pub, caPriv) + if err != nil { + return nil, nil, err + } + cert, err := x509.ParseCertificate(certBytes) + if err != nil { + return nil, nil, err + } + return cert, priv, nil +} + +// GenerateTLSConfigWithLongCertChain generates a tls.Config that uses a long certificate chain. +// The Root CA used is the same as for the config returned from getTLSConfig(). +func GenerateTLSConfigWithLongCertChain(ca *x509.Certificate, caPrivateKey crypto.PrivateKey) (*tls.Config, error) { + const chainLen = 16 + certTempl := &x509.Certificate{ + SerialNumber: big.NewInt(2019), + Subject: pkix.Name{}, + NotBefore: notBefore, + NotAfter: notAfter, + IsCA: true, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + BasicConstraintsValid: true, + } + + lastCA := ca + lastCAPrivKey := caPrivateKey + _, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, err + } + certs := make([]*x509.Certificate, chainLen) + for i := range chainLen { + caBytes, err := x509.CreateCertificate(rand.Reader, certTempl, lastCA, priv.Public(), lastCAPrivKey) + if err != nil { + return nil, err + } + ca, err := x509.ParseCertificate(caBytes) + if err != nil { + return nil, err + } + certs[i] = ca + lastCA = ca + lastCAPrivKey = priv + } + leafCert, leafPrivateKey, err := GenerateLeafCert(lastCA, lastCAPrivKey) + if err != nil { + return nil, err + } + + rawCerts := make([][]byte, chainLen+1) + for i, cert := range certs { + rawCerts[chainLen-i] = cert.Raw + } + rawCerts[0] = leafCert.Raw + + return &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: rawCerts, + PrivateKey: leafPrivateKey, + }}, + NextProtos: []string{ALPN}, + }, nil +} diff --git a/third_party/quic-go/integrationtests/tools/crypto_test.go b/third_party/quic-go/integrationtests/tools/crypto_test.go new file mode 100644 index 0000000..15551ba --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/crypto_test.go @@ -0,0 +1,99 @@ +package tools + +import ( + "crypto/tls" + "crypto/x509" + "io" + "net" + "testing" + + "github.com/stretchr/testify/require" +) + +type countingConn struct { + net.Conn + BytesReceived int +} + +func (c *countingConn) Read(b []byte) (int, error) { + n, err := c.Conn.Read(b) + c.BytesReceived += n + return n, err +} + +func TestGenerateTLSConfig(t *testing.T) { + ca, caPriv, err := GenerateCA() + require.NoError(t, err) + certPool := x509.NewCertPool() + certPool.AddCert(ca) + clientConf := &tls.Config{ + ServerName: "localhost", + RootCAs: certPool, + } + + t.Run("short chain", func(t *testing.T) { + leaf, leafPriv, err := GenerateLeafCert(ca, caPriv) + require.NoError(t, err) + + serverConf := &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: [][]byte{leaf.Raw}, + PrivateKey: leafPriv, + }}, + } + + bytesReceived := testGenerateTLSConfig(t, serverConf, clientConf) + t.Logf("bytes received: %d", bytesReceived) + require.Less(t, bytesReceived, 2000) + }) + + t.Run("long chain", func(t *testing.T) { + serverConf, err := GenerateTLSConfigWithLongCertChain(ca, caPriv) + require.NoError(t, err) + + bytesReceived := testGenerateTLSConfig(t, serverConf, clientConf) + t.Logf("bytes received: %d", bytesReceived) + require.Greater(t, bytesReceived, 5000) + }) +} + +func testGenerateTLSConfig(t *testing.T, serverConf, clientConf *tls.Config) int { + ln, err := tls.Listen("tcp", "127.0.0.1:0", serverConf) + require.NoError(t, err) + defer ln.Close() + + type result struct { + err error + msg string + } + + resultChan := make(chan result, 1) + go func() { + conn, err := ln.Accept() + if err != nil { + resultChan <- result{err: err} + return + } + defer conn.Close() + msg, err := io.ReadAll(conn) + resultChan <- result{err: err, msg: string(msg)} + }() + + tcpConn, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + defer tcpConn.Close() + countingConn := &countingConn{Conn: tcpConn} + + tlsConn := tls.Client(countingConn, clientConf) + require.NoError(t, tlsConn.Handshake()) + + _, err = tlsConn.Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, tlsConn.Close()) + + res := <-resultChan + require.NoError(t, res.err) + require.Equal(t, "foobar", res.msg) + + return countingConn.BytesReceived +} diff --git a/third_party/quic-go/integrationtests/tools/israce/norace.go b/third_party/quic-go/integrationtests/tools/israce/norace.go new file mode 100644 index 0000000..0bf3b8c --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/israce/norace.go @@ -0,0 +1,6 @@ +//go:build !race + +package israce + +// Enabled reports if the race detector is enabled. +const Enabled = false diff --git a/third_party/quic-go/integrationtests/tools/israce/race.go b/third_party/quic-go/integrationtests/tools/israce/race.go new file mode 100644 index 0000000..3751e5c --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/israce/race.go @@ -0,0 +1,6 @@ +//go:build race + +package israce + +// Enabled reports if the race detector is enabled. +const Enabled = true diff --git a/third_party/quic-go/integrationtests/tools/proxy/proxy.go b/third_party/quic-go/integrationtests/tools/proxy/proxy.go new file mode 100644 index 0000000..85340fc --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/proxy/proxy.go @@ -0,0 +1,372 @@ +package quicproxy + +import ( + "errors" + "fmt" + "net" + "os" + "slices" + "sync" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +// Connection is a UDP connection +type connection struct { + ClientAddr *net.UDPAddr // Address of the client + ServerAddr *net.UDPAddr // Address of the server + + mx sync.Mutex + ServerConn *net.UDPConn // UDP connection to server + + incomingPackets chan packetEntry + + Incoming *queue + Outgoing *queue +} + +func (c *connection) queuePacket(t monotime.Time, b []byte) { + c.incomingPackets <- packetEntry{Time: t, Raw: b} +} + +func (c *connection) SwitchConn(conn *net.UDPConn) { + c.mx.Lock() + defer c.mx.Unlock() + + old := c.ServerConn + old.SetReadDeadline(time.Now()) + c.ServerConn = conn +} + +func (c *connection) GetServerConn() *net.UDPConn { + c.mx.Lock() + defer c.mx.Unlock() + + return c.ServerConn +} + +// Direction is the direction a packet is sent. +type Direction int + +const ( + // DirectionIncoming is the direction from the client to the server. + DirectionIncoming Direction = iota + // DirectionOutgoing is the direction from the server to the client. + DirectionOutgoing + // DirectionBoth is both incoming and outgoing + DirectionBoth +) + +type packetEntry struct { + Time monotime.Time + Raw []byte +} + +type queue struct { + sync.Mutex + + timer *time.Timer + Packets []packetEntry // sorted by the packetEntry.Time +} + +func newQueue() *queue { + // there's no way to initialize a time.Timer that's not running + return &queue{timer: time.NewTimer(24 * time.Hour)} +} + +func (q *queue) Add(e packetEntry) { + q.Lock() + defer q.Unlock() + + if len(q.Packets) == 0 { + q.Packets = append(q.Packets, e) + q.timer.Reset(monotime.Until(e.Time)) + return + } + + // The packets slice is sorted by the packetEntry.Time. + // We only need to insert the packet at the correct position. + idx := slices.IndexFunc(q.Packets, func(p packetEntry) bool { + return p.Time.After(e.Time) + }) + if idx == -1 { + q.Packets = append(q.Packets, e) + } else { + q.Packets = slices.Insert(q.Packets, idx, e) + } + if idx == 0 { + q.timer.Reset(monotime.Until(q.Packets[0].Time)) + } +} + +func (q *queue) Get() []byte { + q.Lock() + raw := q.Packets[0].Raw + q.Packets = q.Packets[1:] + if len(q.Packets) > 0 { + q.timer.Reset(monotime.Until(q.Packets[0].Time)) + } + q.Unlock() + return raw +} + +func (q *queue) Timer() <-chan time.Time { return q.timer.C } + +func (q *queue) Close() { q.timer.Stop() } + +func (d Direction) String() string { + switch d { + case DirectionIncoming: + return "Incoming" + case DirectionOutgoing: + return "Outgoing" + case DirectionBoth: + return "both" + default: + panic("unknown direction") + } +} + +// Is says if one direction matches another direction. +// For example, incoming matches both incoming and both, but not outgoing. +func (d Direction) Is(dir Direction) bool { + if d == DirectionBoth || dir == DirectionBoth { + return true + } + return d == dir +} + +// DropCallback is a callback that determines which packet gets dropped. +type DropCallback func(dir Direction, from, to net.Addr, packet []byte) bool + +// DelayCallback is a callback that determines how much delay to apply to a packet. +type DelayCallback func(dir Direction, from, to net.Addr, packet []byte) time.Duration + +// Proxy is a QUIC proxy that can drop and delay packets. +type Proxy struct { + // Conn is the UDP socket that the proxy listens on for incoming packets from clients. + Conn *net.UDPConn + + // ServerAddr is the address of the server that the proxy forwards packets to. + ServerAddr *net.UDPAddr + + // DropPacket is a callback that determines which packet gets dropped. + DropPacket DropCallback + + // DelayPacket is a callback that determines how much delay to apply to a packet. + DelayPacket DelayCallback + + closeChan chan struct{} + logger utils.Logger + + // mapping from client addresses (as host:port) to connection + mutex sync.Mutex + clientDict map[string]*connection +} + +func (p *Proxy) Start() error { + p.clientDict = make(map[string]*connection) + p.closeChan = make(chan struct{}) + p.logger = utils.DefaultLogger.WithPrefix("proxy") + + if err := p.Conn.SetReadBuffer(protocol.DesiredReceiveBufferSize); err != nil { + return err + } + if err := p.Conn.SetWriteBuffer(protocol.DesiredSendBufferSize); err != nil { + return err + } + + p.logger.Debugf("Starting UDP Proxy %s <-> %s", p.Conn.LocalAddr(), p.ServerAddr) + go p.runProxy() + return nil +} + +// SwitchConn switches the connection for a client, +// identified the address that the client is sending from. +func (p *Proxy) SwitchConn(clientAddr *net.UDPAddr, conn *net.UDPConn) error { + if err := conn.SetReadBuffer(protocol.DesiredReceiveBufferSize); err != nil { + return err + } + if err := conn.SetWriteBuffer(protocol.DesiredSendBufferSize); err != nil { + return err + } + p.mutex.Lock() + defer p.mutex.Unlock() + c, ok := p.clientDict[clientAddr.String()] + if !ok { + return fmt.Errorf("client %s not found", clientAddr) + } + c.SwitchConn(conn) + return nil +} + +// Close stops the UDP Proxy +func (p *Proxy) Close() error { + p.mutex.Lock() + defer p.mutex.Unlock() + + close(p.closeChan) + for _, c := range p.clientDict { + if err := c.GetServerConn().Close(); err != nil { + return err + } + c.Incoming.Close() + c.Outgoing.Close() + } + return nil +} + +// LocalAddr is the address the proxy is listening on. +func (p *Proxy) LocalAddr() net.Addr { return p.Conn.LocalAddr() } + +func (p *Proxy) newConnection(cliAddr *net.UDPAddr) (*connection, error) { + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + if err != nil { + return nil, err + } + if err := conn.SetReadBuffer(protocol.DesiredReceiveBufferSize); err != nil { + return nil, err + } + if err := conn.SetWriteBuffer(protocol.DesiredSendBufferSize); err != nil { + return nil, err + } + return &connection{ + ClientAddr: cliAddr, + ServerAddr: p.ServerAddr, + incomingPackets: make(chan packetEntry, 10), + Incoming: newQueue(), + Outgoing: newQueue(), + ServerConn: conn, + }, nil +} + +// runProxy listens on the proxy address and handles incoming packets. +func (p *Proxy) runProxy() error { + for { + buffer := make([]byte, protocol.MaxPacketBufferSize) + n, cliaddr, err := p.Conn.ReadFromUDP(buffer) + if err != nil { + return err + } + raw := buffer[:n] + + p.mutex.Lock() + conn, ok := p.clientDict[cliaddr.String()] + + if !ok { + conn, err = p.newConnection(cliaddr) + if err != nil { + p.mutex.Unlock() + return err + } + p.clientDict[cliaddr.String()] = conn + go p.runIncomingConnection(conn) + go p.runOutgoingConnection(conn) + } + p.mutex.Unlock() + + if p.DropPacket != nil && p.DropPacket(DirectionIncoming, cliaddr, conn.ServerAddr, raw) { + if p.logger.Debug() { + p.logger.Debugf("dropping incoming packet(%d bytes)", n) + } + continue + } + + var delay time.Duration + if p.DelayPacket != nil { + delay = p.DelayPacket(DirectionIncoming, cliaddr, conn.ServerAddr, raw) + } + if delay == 0 { + if p.logger.Debug() { + p.logger.Debugf("forwarding incoming packet (%d bytes) to %s", len(raw), conn.ServerAddr) + } + if _, err := conn.GetServerConn().WriteTo(raw, conn.ServerAddr); err != nil { + return err + } + } else { + now := monotime.Now() + if p.logger.Debug() { + p.logger.Debugf("delaying incoming packet (%d bytes) to %s by %s", len(raw), conn.ServerAddr, delay) + } + conn.queuePacket(now.Add(delay), raw) + } + } +} + +// runConnection handles packets from server to a single client +func (p *Proxy) runOutgoingConnection(conn *connection) error { + outgoingPackets := make(chan packetEntry, 10) + go func() { + for { + buffer := make([]byte, protocol.MaxPacketBufferSize) + n, addr, err := conn.GetServerConn().ReadFrom(buffer) + if err != nil { + // when the connection is switched out, we set a deadline on the old connection, + // in order to return it immediately + if errors.Is(err, os.ErrDeadlineExceeded) { + continue + } + return + } + raw := buffer[0:n] + + if p.DropPacket != nil && p.DropPacket(DirectionOutgoing, addr, conn.ClientAddr, raw) { + if p.logger.Debug() { + p.logger.Debugf("dropping outgoing packet(%d bytes)", n) + } + continue + } + + var delay time.Duration + if p.DelayPacket != nil { + delay = p.DelayPacket(DirectionOutgoing, addr, conn.ClientAddr, raw) + } + if delay == 0 { + if p.logger.Debug() { + p.logger.Debugf("forwarding outgoing packet (%d bytes) to %s", len(raw), conn.ClientAddr) + } + if _, err := p.Conn.WriteToUDP(raw, conn.ClientAddr); err != nil { + return + } + } else { + now := monotime.Now() + if p.logger.Debug() { + p.logger.Debugf("delaying outgoing packet (%d bytes) to %s by %s", len(raw), conn.ClientAddr, delay) + } + outgoingPackets <- packetEntry{Time: now.Add(delay), Raw: raw} + } + } + }() + + for { + select { + case <-p.closeChan: + return nil + case e := <-outgoingPackets: + conn.Outgoing.Add(e) + case <-conn.Outgoing.Timer(): + if _, err := p.Conn.WriteTo(conn.Outgoing.Get(), conn.ClientAddr); err != nil { + return err + } + } + } +} + +func (p *Proxy) runIncomingConnection(conn *connection) error { + for { + select { + case <-p.closeChan: + return nil + case e := <-conn.incomingPackets: + // Send the packet to the server + conn.Incoming.Add(e) + case <-conn.Incoming.Timer(): + if _, err := conn.GetServerConn().WriteTo(conn.Incoming.Get(), conn.ServerAddr); err != nil { + return err + } + } + } +} diff --git a/third_party/quic-go/integrationtests/tools/proxy/proxy_test.go b/third_party/quic-go/integrationtests/tools/proxy/proxy_test.go new file mode 100644 index 0000000..433b101 --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/proxy/proxy_test.go @@ -0,0 +1,503 @@ +package quicproxy + +import ( + "net" + "strconv" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestPacketQueue(t *testing.T) { + q := newQueue() + + getPackets := func() []string { + packets := make([]string, 0, len(q.Packets)) + for _, p := range q.Packets { + packets = append(packets, string(p.Raw)) + } + return packets + } + + require.Empty(t, getPackets()) + now := monotime.Now() + + q.Add(packetEntry{Time: now, Raw: []byte("p3")}) + require.Equal(t, []string{"p3"}, getPackets()) + q.Add(packetEntry{Time: now.Add(time.Second), Raw: []byte("p4")}) + require.Equal(t, []string{"p3", "p4"}, getPackets()) + q.Add(packetEntry{Time: now.Add(-time.Second), Raw: []byte("p1")}) + require.Equal(t, []string{"p1", "p3", "p4"}, getPackets()) + q.Add(packetEntry{Time: now.Add(time.Second), Raw: []byte("p5")}) + require.Equal(t, []string{"p1", "p3", "p4", "p5"}, getPackets()) + q.Add(packetEntry{Time: now.Add(-time.Second), Raw: []byte("p2")}) + require.Equal(t, []string{"p1", "p2", "p3", "p4", "p5"}, getPackets()) +} + +func newUPDConnLocalhost(t testing.TB) *net.UDPConn { + t.Helper() + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + t.Cleanup(func() { conn.Close() }) + return conn +} + +func makePacket(t *testing.T, p protocol.PacketNumber, payload []byte) []byte { + t.Helper() + hdr := wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + Version: protocol.Version1, + Length: 4 + protocol.ByteCount(len(payload)), + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0, 0, 0x13, 0x37}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0, 0, 0x13, 0x37}), + }, + PacketNumber: p, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := hdr.Append(nil, protocol.Version1) + require.NoError(t, err) + b = append(b, payload...) + return b +} + +func readPacketNumber(t *testing.T, b []byte) protocol.PacketNumber { + t.Helper() + hdr, data, _, err := wire.ParsePacket(b) + require.NoError(t, err) + require.Equal(t, protocol.PacketTypeInitial, hdr.Type) + extHdr, err := hdr.ParseExtended(data) + require.NoError(t, err) + return extHdr.PacketNumber +} + +// Set up a dumb UDP server. +// In production this would be a QUIC server. +func runServer(t *testing.T) (*net.UDPAddr, chan []byte) { + done := make(chan struct{}) + t.Cleanup(func() { + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + }) + + serverConn := newUPDConnLocalhost(t) + serverReceivedPackets := make(chan []byte, 100) + go func() { + defer close(done) + for { + buf := make([]byte, protocol.MaxPacketBufferSize) + // the ReadFromUDP will error as soon as the UDP conn is closed + n, addr, err := serverConn.ReadFromUDP(buf) + if err != nil { + return + } + serverReceivedPackets <- buf[:n] + // echo the packet + if _, err := serverConn.WriteToUDP(buf[:n], addr); err != nil { + return + } + } + }() + + return serverConn.LocalAddr().(*net.UDPAddr), serverReceivedPackets +} + +func TestProxyingBackAndForth(t *testing.T) { + serverAddr, _ := runServer(t) + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + // send the first packet + _, err = clientConn.Write(makePacket(t, 1, []byte("foobar"))) + require.NoError(t, err) + // send the second packet + _, err = clientConn.Write(makePacket(t, 2, []byte("decafbad"))) + require.NoError(t, err) + + buf := make([]byte, 1024) + n, err := clientConn.Read(buf) + require.NoError(t, err) + require.Contains(t, string(buf[:n]), "foobar") + n, err = clientConn.Read(buf) + require.NoError(t, err) + require.Contains(t, string(buf[:n]), "decafbad") +} + +func TestDropIncomingPackets(t *testing.T) { + const numPackets = 6 + serverAddr, serverReceivedPackets := runServer(t) + var counter atomic.Int32 + var fromAddr, toAddr atomic.Pointer[net.Addr] + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + DropPacket: func(d Direction, from, to net.Addr, _ []byte) bool { + if d != DirectionIncoming { + return false + } + fromAddr.Store(&from) + toAddr.Store(&to) + return counter.Add(1)%2 == 1 + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + for i := 1; i <= numPackets; i++ { + _, err := clientConn.Write(makePacket(t, protocol.PacketNumber(i), []byte("foobar"+strconv.Itoa(i)))) + require.NoError(t, err) + } + + for range numPackets / 2 { + select { + case <-serverReceivedPackets: + case <-time.After(time.Second): + t.Fatalf("timeout") + } + } + select { + case <-serverReceivedPackets: + t.Fatalf("received unexpected packet") + case <-time.After(100 * time.Millisecond): + } + + require.Equal(t, *fromAddr.Load(), clientConn.LocalAddr()) + require.Equal(t, *toAddr.Load(), serverAddr) +} + +func TestDropOutgoingPackets(t *testing.T) { + const numPackets = 6 + serverAddr, serverReceivedPackets := runServer(t) + var counter atomic.Int32 + var fromAddr, toAddr atomic.Pointer[net.Addr] + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + DropPacket: func(d Direction, from, to net.Addr, _ []byte) bool { + if d != DirectionOutgoing { + return false + } + fromAddr.Store(&from) + toAddr.Store(&to) + return counter.Add(1)%2 == 1 + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + clientReceivedPackets := make(chan struct{}, numPackets) + // receive the packets echoed by the server on client side + go func() { + for { + buf := make([]byte, protocol.MaxPacketBufferSize) + if _, _, err := clientConn.ReadFromUDP(buf); err != nil { + return + } + clientReceivedPackets <- struct{}{} + } + }() + + for i := 1; i <= numPackets; i++ { + _, err := clientConn.Write(makePacket(t, protocol.PacketNumber(i), []byte("foobar"+strconv.Itoa(i)))) + require.NoError(t, err) + } + + for range numPackets / 2 { + select { + case <-clientReceivedPackets: + case <-time.After(time.Second): + t.Fatalf("timeout") + } + } + select { + case <-clientReceivedPackets: + t.Fatalf("received unexpected packet") + case <-time.After(100 * time.Millisecond): + } + require.Len(t, serverReceivedPackets, numPackets) + + require.Equal(t, *fromAddr.Load(), serverAddr) + require.Equal(t, *toAddr.Load(), clientConn.LocalAddr()) +} + +func TestDelayIncomingPackets(t *testing.T) { + const numPackets = 3 + const delay = 200 * time.Millisecond + serverAddr, serverReceivedPackets := runServer(t) + var counter atomic.Int32 + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + DelayPacket: func(d Direction, _, _ net.Addr, _ []byte) time.Duration { + // delay packet 1 by 200 ms + // delay packet 2 by 400 ms + // ... + if d == DirectionOutgoing { + return 0 + } + p := counter.Add(1) + return time.Duration(p) * delay + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + start := time.Now() + for i := 1; i <= numPackets; i++ { + _, err := clientConn.Write(makePacket(t, protocol.PacketNumber(i), []byte("foobar"+strconv.Itoa(i)))) + require.NoError(t, err) + } + + for i := 1; i <= numPackets; i++ { + select { + case data := <-serverReceivedPackets: + require.WithinDuration(t, start.Add(time.Duration(i)*delay), time.Now(), delay/2) + require.Equal(t, protocol.PacketNumber(i), readPacketNumber(t, data)) + case <-time.After(time.Second): + t.Fatalf("timeout waiting for packet %d", i) + } + } +} + +func TestPacketReordering(t *testing.T) { + const delay = 200 * time.Millisecond + expectDelay := func(startTime time.Time, numRTTs int) { + expectedReceiveTime := startTime.Add(time.Duration(numRTTs) * delay) + now := time.Now() + require.True(t, now.After(expectedReceiveTime) || now.Equal(expectedReceiveTime)) + require.True(t, now.Before(expectedReceiveTime.Add(delay/2))) + } + + serverAddr, serverReceivedPackets := runServer(t) + var counter atomic.Int32 + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + DelayPacket: func(d Direction, _, _ net.Addr, _ []byte) time.Duration { + // delay packet 1 by 600 ms + // delay packet 2 by 400 ms + // delay packet 3 by 200 ms + if d == DirectionOutgoing { + return 0 + } + p := counter.Add(1) + return 600*time.Millisecond - time.Duration(p-1)*delay + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + // send 3 packets + start := time.Now() + for i := 1; i <= 3; i++ { + _, err := clientConn.Write(makePacket(t, protocol.PacketNumber(i), []byte("foobar"+strconv.Itoa(i)))) + require.NoError(t, err) + } + for i := 1; i <= 3; i++ { + select { + case packet := <-serverReceivedPackets: + expectDelay(start, i) + expectedPacketNumber := protocol.PacketNumber(4 - i) // 3, 2, 1 in reverse order + require.Equal(t, expectedPacketNumber, readPacketNumber(t, packet)) + case <-time.After(time.Second): + t.Fatalf("timeout waiting for packet %d", i) + } + } +} + +func TestConstantDelay(t *testing.T) { // no reordering expected here + serverAddr, serverReceivedPackets := runServer(t) + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + DelayPacket: func(d Direction, _, _ net.Addr, _ []byte) time.Duration { + if d == DirectionOutgoing { + return 0 + } + return 100 * time.Millisecond + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + // send 100 packets + for i := range 100 { + _, err := clientConn.Write(makePacket(t, protocol.PacketNumber(i), []byte("foobar"+strconv.Itoa(i)))) + require.NoError(t, err) + } + require.Eventually(t, func() bool { return len(serverReceivedPackets) == 100 }, 5*time.Second, 10*time.Millisecond) + timeout := time.After(5 * time.Second) + for i := range 100 { + select { + case packet := <-serverReceivedPackets: + require.Equal(t, protocol.PacketNumber(i), readPacketNumber(t, packet)) + case <-timeout: + t.Fatalf("timeout waiting for packet %d", i) + } + } +} + +func TestDelayOutgoingPackets(t *testing.T) { + const numPackets = 3 + const delay = 200 * time.Millisecond + + serverAddr, serverReceivedPackets := runServer(t) + var counter atomic.Int32 + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverAddr, + DelayPacket: func(d Direction, _, _ net.Addr, _ []byte) time.Duration { + // delay packet 1 by 200 ms + // delay packet 2 by 400 ms + // ... + if d == DirectionIncoming { + return 0 + } + p := counter.Add(1) + return time.Duration(p) * delay + }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + clientConn, err := net.DialUDP("udp", nil, proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + clientReceivedPackets := make(chan []byte, numPackets) + // receive the packets echoed by the server on client side + go func() { + for { + buf := make([]byte, protocol.MaxPacketBufferSize) + n, _, err := clientConn.ReadFromUDP(buf) + if err != nil { + return + } + clientReceivedPackets <- buf[:n] + } + }() + + start := time.Now() + for i := 1; i <= numPackets; i++ { + _, err := clientConn.Write(makePacket(t, protocol.PacketNumber(i), []byte("foobar"+strconv.Itoa(i)))) + require.NoError(t, err) + } + // the packets should have arrived immediately at the server + for range numPackets { + select { + case <-serverReceivedPackets: + case <-time.After(time.Second): + t.Fatalf("timeout") + } + } + require.WithinDuration(t, start, time.Now(), delay/2) + + for i := 1; i <= numPackets; i++ { + select { + case packet := <-clientReceivedPackets: + require.Equal(t, protocol.PacketNumber(i), readPacketNumber(t, packet)) + require.WithinDuration(t, start.Add(time.Duration(i)*delay), time.Now(), delay/2) + case <-time.After(time.Second): + t.Fatalf("timeout waiting for packet %d", i) + } + } +} + +func TestProxySwitchConn(t *testing.T) { + serverConn := newUPDConnLocalhost(t) + + type packet struct { + Data []byte + Addr *net.UDPAddr + } + serverReceivedPackets := make(chan packet, 1) + go func() { + for { + buf := make([]byte, 1000) + n, addr, err := serverConn.ReadFromUDP(buf) + if err != nil { + return + } + serverReceivedPackets <- packet{Data: buf[:n], Addr: addr} + } + }() + + proxy := Proxy{ + Conn: newUPDConnLocalhost(t), + ServerAddr: serverConn.LocalAddr().(*net.UDPAddr), + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + clientConn := newUPDConnLocalhost(t) + _, err := clientConn.WriteToUDP([]byte("hello"), proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + clientConn.SetReadDeadline(time.Now().Add(time.Second)) + + var firstConnAddr *net.UDPAddr + select { + case p := <-serverReceivedPackets: + require.Equal(t, "hello", string(p.Data)) + require.NotEqual(t, clientConn.LocalAddr(), p.Addr) + firstConnAddr = p.Addr + case <-time.After(time.Second): + t.Fatalf("timeout") + } + + _, err = serverConn.WriteToUDP([]byte("hi"), firstConnAddr) + require.NoError(t, err) + buf := make([]byte, 1000) + n, addr, err := clientConn.ReadFromUDP(buf) + require.NoError(t, err) + require.Equal(t, "hi", string(buf[:n])) + require.Equal(t, proxy.LocalAddr(), addr) + + newConn := newUPDConnLocalhost(t) + require.NoError(t, proxy.SwitchConn(clientConn.LocalAddr().(*net.UDPAddr), newConn)) + + _, err = clientConn.WriteToUDP([]byte("foobar"), proxy.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + + select { + case p := <-serverReceivedPackets: + require.Equal(t, "foobar", string(p.Data)) + require.NotEqual(t, clientConn.LocalAddr(), p.Addr) + require.NotEqual(t, firstConnAddr, p.Addr) + require.Equal(t, newConn.LocalAddr(), p.Addr) + case <-time.After(time.Second): + t.Fatalf("timeout") + } + + // the old connection doesn't deliver any packets to the client anymore + _, err = serverConn.WriteTo([]byte("invalid"), firstConnAddr) + require.NoError(t, err) + _, err = serverConn.WriteTo([]byte("foobaz"), newConn.LocalAddr()) + require.NoError(t, err) + n, addr, err = clientConn.ReadFromUDP(buf) + require.NoError(t, err) + require.Equal(t, "foobaz", string(buf[:n])) // "invalid" is not delivered + require.Equal(t, proxy.LocalAddr(), addr) +} diff --git a/third_party/quic-go/integrationtests/tools/qlog.go b/third_party/quic-go/integrationtests/tools/qlog.go new file mode 100644 index 0000000..4870fd4 --- /dev/null +++ b/third_party/quic-go/integrationtests/tools/qlog.go @@ -0,0 +1,55 @@ +package tools + +import ( + "bufio" + "context" + "fmt" + "io" + "log" + "os" + "time" + + "github.com/apernet/quic-go" + h3qlog "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +func QlogTracer(logger io.Writer) qlogwriter.Trace { + filename := fmt.Sprintf("log_%s_transport.qlog", time.Now().Format("2006-01-02T15:04:05")) + fmt.Fprintf(logger, "Creating %s.\n", filename) + f, err := os.Create(filename) + if err != nil { + log.Fatalf("failed to create qlog file: %s", err) + return nil + } + bw := bufio.NewWriter(f) + fileSeq := qlogwriter.NewFileSeq(utils.NewBufferedWriteCloser(bw, f)) + go fileSeq.Run() + return fileSeq +} + +func NewQlogConnectionTracer(logger io.Writer) func(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + return func(_ context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + pers := "server" + if isClient { + pers = "client" + } + filename := fmt.Sprintf("log_%s_%s.qlog", connID, pers) + fmt.Fprintf(logger, "Creating %s.\n", filename) + f, err := os.Create(filename) + if err != nil { + log.Fatalf("failed to create qlog file: %s", err) + return nil + } + fileSeq := qlogwriter.NewConnectionFileSeq( + utils.NewBufferedWriteCloser(bufio.NewWriter(f), f), + isClient, + connID, + []string{qlog.EventSchema, h3qlog.EventSchema}, + ) + go fileSeq.Run() + return fileSeq + } +} diff --git a/third_party/quic-go/integrationtests/versionnegotiation/handshake_test.go b/third_party/quic-go/integrationtests/versionnegotiation/handshake_test.go new file mode 100644 index 0000000..72ba78c --- /dev/null +++ b/third_party/quic-go/integrationtests/versionnegotiation/handshake_test.go @@ -0,0 +1,192 @@ +package versionnegotiation + +import ( + "context" + "errors" + "fmt" + "net" + "testing" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func TestServerSupportsMoreVersionsThanClient(t *testing.T) { + supportedVersions := append([]quic.Version{}, protocol.SupportedVersions...) + protocol.SupportedVersions = append(protocol.SupportedVersions, []protocol.Version{7, 8, 9, 10}...) + defer func() { protocol.SupportedVersions = supportedVersions }() + + var serverEventTracer events.Recorder + serverConfig := &quic.Config{ + Versions: []protocol.Version{7, 8, protocol.SupportedVersions[0], 9}, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: &serverEventTracer} + }, + } + server, err := quic.ListenAddr("localhost:0", getTLSConfig(), serverConfig) + require.NoError(t, err) + defer server.Close() + + var clientEventTracer events.Recorder + conn, err := quic.DialAddr( + context.Background(), + fmt.Sprintf("localhost:%d", server.Addr().(*net.UDPAddr).Port), + getTLSClientConfig(), + maybeAddQLOGTracer(&quic.Config{Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: &clientEventTracer} + }}), + ) + require.NoError(t, err) + + expectedVersion := protocol.SupportedVersions[0] + sconn, err := server.Accept(context.Background()) + require.NoError(t, err) + require.Equal(t, expectedVersion, sconn.ConnectionState().Version) + + require.Equal(t, expectedVersion, conn.ConnectionState().Version) + require.NoError(t, conn.CloseWithError(0, "")) + + select { + case <-sconn.Context().Done(): + // Expected behavior + case <-time.After(5 * time.Second): + t.Fatal("Timeout waiting for connection to close") + } + + require.Empty(t, clientEventTracer.Events(qlog.VersionNegotiationReceived{})) + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionInformation{ + ClientVersions: protocol.SupportedVersions, + ChosenVersion: expectedVersion, + }, + }, + clientEventTracer.Events(qlog.VersionInformation{}), + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionInformation{ + ServerVersions: serverConfig.Versions, + ChosenVersion: expectedVersion, + }, + }, + serverEventTracer.Events(qlog.VersionInformation{}), + ) +} + +func TestClientSupportsMoreVersionsThanServer(t *testing.T) { + supportedVersions := append([]quic.Version{}, protocol.SupportedVersions...) + protocol.SupportedVersions = append(protocol.SupportedVersions, []protocol.Version{7, 8, 9, 10}...) + defer func() { protocol.SupportedVersions = supportedVersions }() + + expectedVersion := protocol.SupportedVersions[0] + // The server doesn't support the highest supported version, which is the first one the client will try, + // but it supports a bunch of versions that the client doesn't speak + var serverEventTracer events.Recorder + serverConfig := &quic.Config{ + Versions: supportedVersions, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: &serverEventTracer} + }, + } + server, err := quic.ListenAddr("localhost:0", getTLSConfig(), serverConfig) + require.NoError(t, err) + defer server.Close() + + clientVersions := []protocol.Version{7, 8, 9, protocol.SupportedVersions[0], 10} + var clientEventTracer events.Recorder + conn, err := quic.DialAddr( + context.Background(), + fmt.Sprintf("localhost:%d", server.Addr().(*net.UDPAddr).Port), + getTLSClientConfig(), + maybeAddQLOGTracer(&quic.Config{ + Versions: clientVersions, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: &clientEventTracer} + }, + }), + ) + require.NoError(t, err) + + sconn, err := server.Accept(context.Background()) + require.NoError(t, err) + require.Equal(t, expectedVersion, sconn.ConnectionState().Version) + + require.Equal(t, protocol.SupportedVersions[0], conn.ConnectionState().Version) + require.NoError(t, conn.CloseWithError(0, "")) + + select { + case <-sconn.Context().Done(): + // Expected behavior + case <-time.After(5 * time.Second): + t.Fatal("Timeout waiting for connection to close") + } + + require.Len(t, clientEventTracer.Events(qlog.VersionNegotiationReceived{}), 1) + supportedVersionInclGreased := clientEventTracer.Events(qlog.VersionNegotiationReceived{})[0].(qlog.VersionNegotiationReceived).SupportedVersions + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionInformation{ + ClientVersions: clientVersions, + ServerVersions: supportedVersionInclGreased, + ChosenVersion: expectedVersion, + }, + }, + clientEventTracer.Events(qlog.VersionInformation{}), + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionInformation{ + ServerVersions: supportedVersions, + ChosenVersion: expectedVersion, + }, + }, + serverEventTracer.Events(qlog.VersionInformation{}), + ) +} + +func TestServerDisablesVersionNegotiation(t *testing.T) { + // The server doesn't support the highest supported version, which is the first one the client will try, + // but it supports a bunch of versions that the client doesn't speak + var serverEventTracer events.Recorder + serverConfig := &quic.Config{ + Versions: []protocol.Version{quic.Version1}, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: &serverEventTracer} + }, + } + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + require.NoError(t, err) + tr := &quic.Transport{ + Conn: conn, + DisableVersionNegotiationPackets: true, + } + ln, err := tr.Listen(getTLSConfig(), serverConfig) + require.NoError(t, err) + defer ln.Close() + + var clientEventTracer events.Recorder + _, err = quic.DialAddr( + context.Background(), + fmt.Sprintf("localhost:%d", conn.LocalAddr().(*net.UDPAddr).Port), + getTLSClientConfig(), + maybeAddQLOGTracer(&quic.Config{ + Versions: []protocol.Version{quic.Version2}, + Tracer: func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace { + return &events.Trace{Recorder: &clientEventTracer} + }, + HandshakeIdleTimeout: 100 * time.Millisecond, + }), + ) + require.Error(t, err) + var nerr net.Error + require.True(t, errors.As(err, &nerr)) + require.True(t, nerr.Timeout()) + require.Empty(t, clientEventTracer.Events(qlog.VersionNegotiationReceived{})) +} diff --git a/third_party/quic-go/integrationtests/versionnegotiation/rtt_test.go b/third_party/quic-go/integrationtests/versionnegotiation/rtt_test.go new file mode 100644 index 0000000..408525f --- /dev/null +++ b/third_party/quic-go/integrationtests/versionnegotiation/rtt_test.go @@ -0,0 +1,58 @@ +package versionnegotiation + +import ( + "context" + "net" + "testing" + "time" + + "github.com/apernet/quic-go" + quicproxy "github.com/apernet/quic-go/integrationtests/tools/proxy" + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +const rtt = 400 * time.Millisecond + +func expectDurationInRTTs(t *testing.T, startTime time.Time, num int) { + t.Helper() + testDuration := time.Since(startTime) + rtts := float32(testDuration) / float32(rtt) + require.GreaterOrEqual(t, rtts, float32(num)) + require.Less(t, rtts, float32(num+1)) +} + +func TestVersionNegotiationFailure(t *testing.T) { + if len(protocol.SupportedVersions) == 1 { + t.Fatal("Test requires at least 2 supported versions.") + } + + serverConfig := &quic.Config{} + serverConfig.Versions = protocol.SupportedVersions[:1] + ln, err := quic.ListenAddr("localhost:0", getTLSConfig(), serverConfig) + require.NoError(t, err) + defer ln.Close() + + proxyConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + defer proxyConn.Close() + // start the proxy + proxy := quicproxy.Proxy{ + Conn: proxyConn, + ServerAddr: ln.Addr().(*net.UDPAddr), + DelayPacket: func(quicproxy.Direction, net.Addr, net.Addr, []byte) time.Duration { return rtt / 2 }, + } + require.NoError(t, proxy.Start()) + defer proxy.Close() + + startTime := time.Now() + _, err = quic.DialAddr( + context.Background(), + proxy.LocalAddr().String(), + getTLSClientConfig(), + maybeAddQLOGTracer(&quic.Config{Versions: protocol.SupportedVersions[1:2]}), + ) + require.Error(t, err) + expectDurationInRTTs(t, startTime, 1) +} diff --git a/third_party/quic-go/integrationtests/versionnegotiation/test_helper_test.go b/third_party/quic-go/integrationtests/versionnegotiation/test_helper_test.go new file mode 100644 index 0000000..e27e134 --- /dev/null +++ b/third_party/quic-go/integrationtests/versionnegotiation/test_helper_test.go @@ -0,0 +1,111 @@ +package versionnegotiation + +import ( + "context" + "crypto/tls" + "crypto/x509" + "flag" + "os" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/integrationtests/tools" + "github.com/apernet/quic-go/qlogwriter" +) + +var ( + enableQlog bool + tlsConfig *tls.Config + tlsClientConfig *tls.Config +) + +func init() { + flag.BoolVar(&enableQlog, "qlog", false, "enable qlog") + + ca, caPrivateKey, err := tools.GenerateCA() + if err != nil { + panic(err) + } + leafCert, leafPrivateKey, err := tools.GenerateLeafCert(ca, caPrivateKey) + if err != nil { + panic(err) + } + tlsConfig = &tls.Config{ + Certificates: []tls.Certificate{{ + Certificate: [][]byte{leafCert.Raw}, + PrivateKey: leafPrivateKey, + }}, + NextProtos: []string{tools.ALPN}, + } + + root := x509.NewCertPool() + root.AddCert(ca) + tlsClientConfig = &tls.Config{ + ServerName: "localhost", + RootCAs: root, + NextProtos: []string{tools.ALPN}, + } +} + +func getTLSConfig() *tls.Config { return tlsConfig } +func getTLSClientConfig() *tls.Config { return tlsClientConfig } + +type multiplexedRecorder struct { + Recorders []qlogwriter.Recorder +} + +var _ qlogwriter.Recorder = &multiplexedRecorder{} + +func (r *multiplexedRecorder) Close() error { + for _, recorder := range r.Recorders { + recorder.Close() + } + return nil +} + +func (r *multiplexedRecorder) RecordEvent(ev qlogwriter.Event) { + for _, recorder := range r.Recorders { + recorder.RecordEvent(ev) + } +} + +type multiplexedTrace struct { + Traces []qlogwriter.Trace +} + +var _ qlogwriter.Trace = &multiplexedTrace{} + +func (t *multiplexedTrace) SupportsSchemas(schema string) bool { return true } + +func (t *multiplexedTrace) AddProducer() qlogwriter.Recorder { + recorders := make([]qlogwriter.Recorder, 0, len(t.Traces)) + for _, tr := range t.Traces { + recorders = append(recorders, tr.AddProducer()) + } + return &multiplexedRecorder{Recorders: recorders} +} + +func maybeAddQLOGTracer(c *quic.Config) *quic.Config { + if c == nil { + c = &quic.Config{} + } + if !enableQlog { + return c + } + qlogger := tools.NewQlogConnectionTracer(os.Stdout) + if c.Tracer == nil { + c.Tracer = qlogger + } else if qlogger != nil { + origTracer := c.Tracer + c.Tracer = func(ctx context.Context, p bool, connID quic.ConnectionID) qlogwriter.Trace { + var traces []qlogwriter.Trace + if origTracer != nil { + traces = append(traces, origTracer(ctx, p, connID)) + } + if qlogger != nil { + traces = append(traces, qlogger(ctx, p, connID)) + } + return &multiplexedTrace{Traces: traces} + } + } + return c +} diff --git a/third_party/quic-go/interface.go b/third_party/quic-go/interface.go new file mode 100644 index 0000000..0e6ef58 --- /dev/null +++ b/third_party/quic-go/interface.go @@ -0,0 +1,243 @@ +package quic + +import ( + "context" + "crypto/tls" + "errors" + "net" + "slices" + "time" + + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter" +) + +// The StreamID is the ID of a QUIC stream. +type StreamID = protocol.StreamID + +// A Version is a QUIC version number. +type Version = protocol.Version + +const ( + // Version1 is RFC 9000 + Version1 = protocol.Version1 + // Version2 is RFC 9369 + Version2 = protocol.Version2 +) + +// SupportedVersions returns the support versions, sorted in descending order of preference. +func SupportedVersions() []Version { + // clone the slice to prevent the caller from modifying the slice + return slices.Clone(protocol.SupportedVersions) +} + +// A ClientToken is a token received by the client. +// It can be used to skip address validation on future connection attempts. +type ClientToken struct { + data []byte + rtt time.Duration +} + +type TokenStore interface { + // Pop searches for a ClientToken associated with the given key. + // Since tokens are not supposed to be reused, it must remove the token from the cache. + // It returns nil when no token is found. + Pop(key string) (token *ClientToken) + + // Put adds a token to the cache with the given key. It might get called + // multiple times in a connection. + Put(key string, token *ClientToken) +} + +// Err0RTTRejected is the returned from: +// - Open{Uni}Stream{Sync} +// - Accept{Uni}Stream +// - Stream.Read and Stream.Write +// +// when the server rejects a 0-RTT connection attempt. +var Err0RTTRejected = errors.New("0-RTT rejected") + +// ErrWouldBlock is returned by [SendStream.TryWriteAll] if the entire slice can't be queued immediately. +var ErrWouldBlock = errors.New("operation would block") + +// ErrWriteLimitReached is returned by [SendStream.WriteWithLimit] when its limiter prevents accepting the entire slice. +var ErrWriteLimitReached = errors.New("write limit reached") + +// QUICVersionContextKey can be used to find out the QUIC version of a TLS handshake from the +// context returned by tls.Config.ClientInfo.Context. +var QUICVersionContextKey = handshake.QUICVersionContextKey + +// StatelessResetKey is a key used to derive stateless reset tokens. +type StatelessResetKey [32]byte + +// TokenGeneratorKey is a key used to encrypt session resumption tokens. +type TokenGeneratorKey = handshake.TokenProtectorKey + +// A ConnectionID is a QUIC Connection ID, as defined in RFC 9000. +// It is not able to handle QUIC Connection IDs longer than 20 bytes, +// as they are allowed by RFC 8999. +type ConnectionID = protocol.ConnectionID + +// ConnectionIDFromBytes interprets b as a [ConnectionID]. It panics if b is +// longer than 20 bytes. +func ConnectionIDFromBytes(b []byte) ConnectionID { + return protocol.ParseConnectionID(b) +} + +// A ConnectionIDGenerator allows the application to take control over the generation of Connection IDs. +// Connection IDs generated by an implementation must be of constant length. +type ConnectionIDGenerator interface { + // GenerateConnectionID generates a new Connection ID. + // Generated Connection IDs must be unique and observers should not be able to correlate two Connection IDs. + GenerateConnectionID() (ConnectionID, error) + + // ConnectionIDLen returns the length of Connection IDs generated by this implementation. + // Implementations must return constant-length Connection IDs with lengths between 0 and 20 bytes. + // A length of 0 can only be used when an endpoint doesn't need to multiplex connections during migration. + ConnectionIDLen() int +} + +// Config contains all configuration data needed for a QUIC server or client. +type Config struct { + // GetConfigForClient is called for incoming connections. + // If the error is not nil, the connection attempt is refused. + GetConfigForClient func(info *ClientInfo) (*Config, error) + // The QUIC versions that can be negotiated. + // If not set, it uses all versions available. + Versions []Version + // HandshakeIdleTimeout is the idle timeout before completion of the handshake. + // If we don't receive any packet from the peer within this time, the connection attempt is aborted. + // Additionally, if the handshake doesn't complete in twice this time, the connection attempt is also aborted. + // If this value is zero, the timeout is set to 5 seconds. + HandshakeIdleTimeout time.Duration + // MaxIdleTimeout is the maximum duration that may pass without any incoming network activity. + // The actual value for the idle timeout is the minimum of this value and the peer's. + // This value only applies after the handshake has completed. + // If the timeout is exceeded, the connection is closed. + // If this value is zero, the timeout is set to 30 seconds. + MaxIdleTimeout time.Duration + // The TokenStore stores tokens received from the server. + // Tokens are used to skip address validation on future connection attempts. + // The key used to store tokens is the ServerName from the tls.Config, if set + // otherwise the token is associated with the server's IP address. + TokenStore TokenStore + // InitialStreamReceiveWindow is the initial size of the stream-level flow control window for receiving data. + // If the application is consuming data quickly enough, the flow control auto-tuning algorithm + // will increase the window up to MaxStreamReceiveWindow. + // If this value is zero, it will default to 512 KB. + // Values larger than the maximum varint (quicvarint.Max) will be clipped to that value. + InitialStreamReceiveWindow uint64 + // MaxStreamReceiveWindow is the maximum stream-level flow control window for receiving data. + // If this value is zero, it will default to 6 MB. + // Values larger than the maximum varint (quicvarint.Max) will be clipped to that value. + MaxStreamReceiveWindow uint64 + // InitialConnectionReceiveWindow is the initial size of the stream-level flow control window for receiving data. + // If the application is consuming data quickly enough, the flow control auto-tuning algorithm + // will increase the window up to MaxConnectionReceiveWindow. + // If this value is zero, it will default to 512 KB. + // Values larger than the maximum varint (quicvarint.Max) will be clipped to that value. + InitialConnectionReceiveWindow uint64 + // MaxConnectionReceiveWindow is the connection-level flow control window for receiving data. + // If this value is zero, it will default to 15 MB. + // Values larger than the maximum varint (quicvarint.Max) will be clipped to that value. + MaxConnectionReceiveWindow uint64 + // AllowConnectionWindowIncrease is called every time the connection flow controller attempts + // to increase the connection flow control window. + // If set, the caller can prevent an increase of the window. Typically, it would do so to + // limit the memory usage. + // To avoid deadlocks, it is not valid to call other functions on the connection or on streams + // in this callback. + AllowConnectionWindowIncrease func(conn *Conn, delta uint64) bool + // MaxIncomingStreams is the maximum number of concurrent bidirectional streams that a peer is allowed to open. + // If not set, it will default to 100. + // If set to a negative value, it doesn't allow any bidirectional streams. + // Values larger than 2^60 will be clipped to that value. + MaxIncomingStreams int64 + // MaxIncomingUniStreams is the maximum number of concurrent unidirectional streams that a peer is allowed to open. + // If not set, it will default to 100. + // If set to a negative value, it doesn't allow any unidirectional streams. + // Values larger than 2^60 will be clipped to that value. + MaxIncomingUniStreams int64 + // KeepAlivePeriod defines whether this peer will periodically send a packet to keep the connection alive. + // If set to 0, then no keep alive is sent. Otherwise, the keep alive is sent on that period (or at most + // every half of MaxIdleTimeout, whichever is smaller). + KeepAlivePeriod time.Duration + // InitialPacketSize is the initial size (and the lower limit) for packets sent. + // Under most circumstances, it is not necessary to manually set this value, + // since path MTU discovery quickly finds the path's MTU. + // If set too high, the path might not support packets of that size, leading to a timeout of the QUIC handshake. + // Values below 1200 are invalid. + InitialPacketSize uint16 + // DisablePathMTUDiscovery disables Path MTU Discovery (RFC 8899). + // This allows the sending of QUIC packets that fully utilize the available MTU of the path. + // Path MTU discovery is only available on systems that allow setting of the Don't Fragment (DF) bit. + DisablePathMTUDiscovery bool + // Allow0RTT allows the application to decide if a 0-RTT connection attempt should be accepted. + // Only valid for the server. + Allow0RTT bool + // Enable QUIC datagram support (RFC 9221). + EnableDatagrams bool + // OmitMaxDatagramFrameSize omits the max_datagram_frame_size transport parameter, + // even when QUIC datagram support is enabled. + OmitMaxDatagramFrameSize bool + // AssumePeerMaxDatagramFrameSize treats peers that omit max_datagram_frame_size + // as supporting DATAGRAM frames up to this size. This is a non-standard extension. + AssumePeerMaxDatagramFrameSize int64 + // Enable QUIC Stream Resets with Partial Delivery. + // See https://datatracker.ietf.org/doc/html/draft-ietf-quic-reliable-stream-reset-09. + EnableStreamResetPartialDelivery bool + + Tracer func(ctx context.Context, isClient bool, connID ConnectionID) qlogwriter.Trace + + MaxDatagramFrameSize int64 + + // DisablePathManager disables path manager. + // for hysteria2 port hopping, direct change remote address without connection migration logic + DisablePathManager bool + + // ChromeParrot makes the client's QUIC handshake look like Google Chrome's. + // It overrides the flow control windows, stream limits, idle timeout and + // packet size with Chrome's values, encodes the transport parameters the way + // Chrome does (see wire.marshalChrome), and applies Chrome's chaos + // protection to the Initial packets. + // + // Client side only; it has no effect on a listener. Because it pins the + // values above, settings that conflict with Chrome's are ignored. + ChromeParrot bool +} + +// ClientInfo contains information about an incoming connection attempt. +type ClientInfo struct { + // RemoteAddr is the remote address on the Initial packet. + // Unless AddrVerified is set, the address is not yet verified, and could be a spoofed IP address. + RemoteAddr net.Addr + // AddrVerified says if the remote address was verified using QUIC's Retry mechanism. + // Note that the Retry mechanism costs one network roundtrip, + // and is not performed unless Transport.MaxUnvalidatedHandshakes is surpassed. + AddrVerified bool +} + +// ConnectionState records basic details about a QUIC connection. +type ConnectionState struct { + // TLS contains information about the TLS connection state, incl. the tls.ConnectionState. + TLS tls.ConnectionState + // SupportsDatagrams indicates support for QUIC datagrams (RFC 9221). + SupportsDatagrams struct { + // Remote is true if the peer advertised datagram support. + // Local is true if datagram support was enabled via Config.EnableDatagrams. + Remote, Local bool + } + // SupportsStreamResetPartialDelivery indicates support for QUIC Stream Resets with Partial Delivery. + SupportsStreamResetPartialDelivery struct { + // Remote is true if the peer advertised support. + // Local is true if support was enabled via Config.EnableStreamResetPartialDelivery. + Remote, Local bool + } + // Used0RTT says if 0-RTT resumption was used. + Used0RTT bool + // Version is the QUIC version of the QUIC connection. + Version Version + // GSO says if generic segmentation offload is used. + GSO bool +} diff --git a/third_party/quic-go/internal/ackhandler/ack_eliciting.go b/third_party/quic-go/internal/ackhandler/ack_eliciting.go new file mode 100644 index 0000000..0e7f4ed --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/ack_eliciting.go @@ -0,0 +1,33 @@ +package ackhandler + +import "github.com/apernet/quic-go/internal/wire" + +// IsFrameTypeAckEliciting returns true if the frame is ack-eliciting. +func IsFrameTypeAckEliciting(t wire.FrameType) bool { + //nolint:exhaustive // The default case catches the rest. + switch t { + case wire.FrameTypeAck, wire.FrameTypeAckECN: + return false + case wire.FrameTypeConnectionClose, wire.FrameTypeApplicationClose: + return false + default: + return true + } +} + +// IsFrameAckEliciting returns true if the frame is ack-eliciting. +func IsFrameAckEliciting(f wire.Frame) bool { + _, isAck := f.(*wire.AckFrame) + _, isConnectionClose := f.(*wire.ConnectionCloseFrame) + return !isAck && !isConnectionClose +} + +// HasAckElicitingFrames returns true if at least one frame is ack-eliciting. +func HasAckElicitingFrames(fs []Frame) bool { + for _, f := range fs { + if IsFrameAckEliciting(f.Frame) { + return true + } + } + return false +} diff --git a/third_party/quic-go/internal/ackhandler/ack_eliciting_test.go b/third_party/quic-go/internal/ackhandler/ack_eliciting_test.go new file mode 100644 index 0000000..973cbfa --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/ack_eliciting_test.go @@ -0,0 +1,74 @@ +package ackhandler + +import ( + "testing" + + "github.com/apernet/quic-go/internal/wire" + "github.com/stretchr/testify/require" +) + +func TestIsFrameTypeAckEliciting(t *testing.T) { + testCases := map[wire.FrameType]bool{ + wire.FrameTypePing: true, + wire.FrameTypeAck: false, + wire.FrameTypeAckECN: false, + wire.FrameTypeResetStream: true, + wire.FrameTypeStopSending: true, + wire.FrameTypeCrypto: true, + wire.FrameTypeNewToken: true, + wire.FrameType(0x08): true, + wire.FrameType(0x09): true, + wire.FrameType(0x0a): true, + wire.FrameType(0x0b): true, + wire.FrameType(0x0c): true, + wire.FrameType(0x0d): true, + wire.FrameType(0x0e): true, + wire.FrameType(0x0f): true, + wire.FrameTypeMaxData: true, + wire.FrameTypeMaxStreamData: true, + wire.FrameTypeBidiMaxStreams: true, + wire.FrameTypeUniMaxStreams: true, + wire.FrameTypeDataBlocked: true, + wire.FrameTypeStreamDataBlocked: true, + wire.FrameTypeBidiStreamBlocked: true, + wire.FrameTypeUniStreamBlocked: true, + wire.FrameTypeNewConnectionID: true, + wire.FrameTypeRetireConnectionID: true, + wire.FrameTypePathChallenge: true, + wire.FrameTypePathResponse: true, + wire.FrameTypeConnectionClose: false, + wire.FrameTypeApplicationClose: false, + wire.FrameTypeHandshakeDone: true, + wire.FrameTypeResetStreamAt: true, + wire.FrameTypeDatagramNoLength: true, + wire.FrameTypeDatagramWithLength: true, + wire.FrameTypeAckFrequency: true, + wire.FrameTypeImmediateAck: true, + } + + for ft, expected := range testCases { + require.Equal(t, expected, IsFrameTypeAckEliciting(ft), "unexpected result for frame type 0x%x", ft) + } +} + +func TestAckElicitingFrames(t *testing.T) { + testCases := map[wire.Frame]bool{ + &wire.AckFrame{}: false, + &wire.ConnectionCloseFrame{}: false, + &wire.DataBlockedFrame{}: true, + &wire.PingFrame{}: true, + &wire.ResetStreamFrame{}: true, + &wire.StreamFrame{}: true, + &wire.DatagramFrame{}: true, + &wire.MaxDataFrame{}: true, + &wire.MaxStreamDataFrame{}: true, + &wire.StopSendingFrame{}: true, + &wire.AckFrequencyFrame{}: true, + &wire.ImmediateAckFrame{}: true, + } + + for f, expected := range testCases { + require.Equal(t, expected, IsFrameAckEliciting(f)) + require.Equal(t, expected, HasAckElicitingFrames([]Frame{{Frame: f}})) + } +} diff --git a/third_party/quic-go/internal/ackhandler/cc_adapter.go b/third_party/quic-go/internal/ackhandler/cc_adapter.go new file mode 100644 index 0000000..d6daa0f --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/cc_adapter.go @@ -0,0 +1,62 @@ +package ackhandler + +import ( + "github.com/apernet/quic-go/congestion" + cgInternal "github.com/apernet/quic-go/internal/congestion" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +var _ cgInternal.SendAlgorithmWithDebugInfos = &ccAdapter{} + +type ccAdapter struct { + CC congestion.CongestionControl +} + +func (a *ccAdapter) TimeUntilSend(bytesInFlight protocol.ByteCount) monotime.Time { + return a.CC.TimeUntilSend(congestion.ByteCount(bytesInFlight)) +} + +func (a *ccAdapter) HasPacingBudget(now monotime.Time) bool { + return a.CC.HasPacingBudget(now) +} + +func (a *ccAdapter) OnPacketSent(sentTime monotime.Time, bytesInFlight protocol.ByteCount, packetNumber protocol.PacketNumber, bytes protocol.ByteCount, isRetransmittable bool) { + a.CC.OnPacketSent(sentTime, congestion.ByteCount(bytesInFlight), congestion.PacketNumber(packetNumber), congestion.ByteCount(bytes), isRetransmittable) +} + +func (a *ccAdapter) CanSend(bytesInFlight protocol.ByteCount) bool { + return a.CC.CanSend(congestion.ByteCount(bytesInFlight)) +} + +func (a *ccAdapter) MaybeExitSlowStart() { + a.CC.MaybeExitSlowStart() +} + +func (a *ccAdapter) OnPacketAcked(number protocol.PacketNumber, ackedBytes protocol.ByteCount, priorInFlight protocol.ByteCount, eventTime monotime.Time) { + a.CC.OnPacketAcked(congestion.PacketNumber(number), congestion.ByteCount(ackedBytes), congestion.ByteCount(priorInFlight), eventTime) +} + +func (a *ccAdapter) OnCongestionEvent(number protocol.PacketNumber, lostBytes protocol.ByteCount, priorInFlight protocol.ByteCount) { + a.CC.OnCongestionEvent(congestion.PacketNumber(number), congestion.ByteCount(lostBytes), congestion.ByteCount(priorInFlight)) +} + +func (a *ccAdapter) OnRetransmissionTimeout(packetsRetransmitted bool) { + a.CC.OnRetransmissionTimeout(packetsRetransmitted) +} + +func (a *ccAdapter) SetMaxDatagramSize(size protocol.ByteCount) { + a.CC.SetMaxDatagramSize(congestion.ByteCount(size)) +} + +func (a *ccAdapter) InSlowStart() bool { + return a.CC.InSlowStart() +} + +func (a *ccAdapter) InRecovery() bool { + return a.CC.InRecovery() +} + +func (a *ccAdapter) GetCongestionWindow() protocol.ByteCount { + return protocol.ByteCount(a.CC.GetCongestionWindow()) +} diff --git a/third_party/quic-go/internal/ackhandler/cc_adapter_ex.go b/third_party/quic-go/internal/ackhandler/cc_adapter_ex.go new file mode 100644 index 0000000..e8d3bf0 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/cc_adapter_ex.go @@ -0,0 +1,69 @@ +package ackhandler + +import ( + "github.com/apernet/quic-go/congestion" + cgInternal "github.com/apernet/quic-go/internal/congestion" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +var ( + _ cgInternal.SendAlgorithmEx = &ccAdapterEx{} + _ cgInternal.SendAlgorithmWithDebugInfos = &ccAdapterEx{} +) + +type ccAdapterEx struct { + CC congestion.CongestionControlEx +} + +func (a *ccAdapterEx) TimeUntilSend(bytesInFlight protocol.ByteCount) monotime.Time { + return a.CC.TimeUntilSend(congestion.ByteCount(bytesInFlight)) +} + +func (a *ccAdapterEx) HasPacingBudget(now monotime.Time) bool { + return a.CC.HasPacingBudget(now) +} + +func (a *ccAdapterEx) OnPacketSent(sentTime monotime.Time, bytesInFlight protocol.ByteCount, packetNumber protocol.PacketNumber, bytes protocol.ByteCount, isRetransmittable bool) { + a.CC.OnPacketSent(sentTime, congestion.ByteCount(bytesInFlight), congestion.PacketNumber(packetNumber), congestion.ByteCount(bytes), isRetransmittable) +} + +func (a *ccAdapterEx) CanSend(bytesInFlight protocol.ByteCount) bool { + return a.CC.CanSend(congestion.ByteCount(bytesInFlight)) +} + +func (a *ccAdapterEx) MaybeExitSlowStart() { + a.CC.MaybeExitSlowStart() +} + +func (a *ccAdapterEx) OnPacketAcked(number protocol.PacketNumber, ackedBytes protocol.ByteCount, priorInFlight protocol.ByteCount, eventTime monotime.Time) { + a.CC.OnPacketAcked(congestion.PacketNumber(number), congestion.ByteCount(ackedBytes), congestion.ByteCount(priorInFlight), eventTime) +} + +func (a *ccAdapterEx) OnCongestionEvent(number protocol.PacketNumber, lostBytes protocol.ByteCount, priorInFlight protocol.ByteCount) { + a.CC.OnCongestionEvent(congestion.PacketNumber(number), congestion.ByteCount(lostBytes), congestion.ByteCount(priorInFlight)) +} + +func (a *ccAdapterEx) OnCongestionEventEx(priorInFlight protocol.ByteCount, eventTime monotime.Time, ackedPackets []congestion.AckedPacketInfo, lostPackets []congestion.LostPacketInfo) { + a.CC.OnCongestionEventEx(congestion.ByteCount(priorInFlight), eventTime, ackedPackets, lostPackets) +} + +func (a *ccAdapterEx) OnRetransmissionTimeout(packetsRetransmitted bool) { + a.CC.OnRetransmissionTimeout(packetsRetransmitted) +} + +func (a *ccAdapterEx) SetMaxDatagramSize(size protocol.ByteCount) { + a.CC.SetMaxDatagramSize(congestion.ByteCount(size)) +} + +func (a *ccAdapterEx) InSlowStart() bool { + return a.CC.InSlowStart() +} + +func (a *ccAdapterEx) InRecovery() bool { + return a.CC.InRecovery() +} + +func (a *ccAdapterEx) GetCongestionWindow() protocol.ByteCount { + return protocol.ByteCount(a.CC.GetCongestionWindow()) +} diff --git a/third_party/quic-go/internal/ackhandler/ecn.go b/third_party/quic-go/internal/ackhandler/ecn.go new file mode 100644 index 0000000..8195f04 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/ecn.go @@ -0,0 +1,340 @@ +package ackhandler + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +type ecnState uint8 + +const ( + ecnStateInitial ecnState = iota + ecnStateTesting + ecnStateUnknown + ecnStateCapable + ecnStateFailed +) + +const ( + // ecnFailedNoECNCounts is emitted when an ACK acknowledges ECN-marked packets, + // but doesn't contain any ECN counts + ecnFailedNoECNCounts = "ACK doesn't contain ECN marks" + // ecnFailedDecreasedECNCounts is emitted when an ACK frame decreases ECN counts + ecnFailedDecreasedECNCounts = "ACK decreases ECN counts" + // ecnFailedLostAllTestingPackets is emitted when all ECN testing packets are declared lost + ecnFailedLostAllTestingPackets = "all ECN testing packets declared lost" + // ecnFailedMoreECNCountsThanSent is emitted when an ACK contains more ECN counts than ECN-marked packets were sent + ecnFailedMoreECNCountsThanSent = "ACK contains more ECN counts than ECN-marked packets sent" + // ecnFailedTooFewECNCounts is emitted when an ACK contains fewer ECN counts than it acknowledges packets + ecnFailedTooFewECNCounts = "ACK contains fewer new ECN counts than acknowledged ECN-marked packets" + // ecnFailedManglingDetected is emitted when the path marks all ECN-marked packets as CE + ecnFailedManglingDetected = "ECN mangling detected" +) + +// must fit into an uint8, otherwise numSentTesting and numLostTesting must have a larger type +const numECNTestingPackets = 10 + +type ecnHandler interface { + SentPacket(protocol.PacketNumber, protocol.ECN) + Mode() protocol.ECN + HandleNewlyAcked(packets []packetWithPacketNumber, ect0, ect1, ecnce int64) (congested bool) + LostPacket(protocol.PacketNumber) +} + +// The ecnTracker performs ECN validation of a path. +// Once failed, it doesn't do any re-validation of the path. +// It is designed only work for 1-RTT packets, it doesn't handle multiple packet number spaces. +// In order to avoid revealing any internal state to on-path observers, +// callers should make sure to start using ECN (i.e. calling Mode) for the very first 1-RTT packet sent. +// The validation logic implemented here strictly follows the algorithm described in RFC 9000 section 13.4.2 and A.4. +type ecnTracker struct { + state ecnState + numSentTesting, numLostTesting uint8 + + firstTestingPacket protocol.PacketNumber + lastTestingPacket protocol.PacketNumber + firstCapablePacket protocol.PacketNumber + + numSentECT0, numSentECT1 int64 + numAckedECT0, numAckedECT1, numAckedECNCE int64 + + qlogger qlogwriter.Recorder + logger utils.Logger +} + +var _ ecnHandler = &ecnTracker{} + +func newECNTracker(logger utils.Logger, qlogger qlogwriter.Recorder) *ecnTracker { + return &ecnTracker{ + firstTestingPacket: protocol.InvalidPacketNumber, + lastTestingPacket: protocol.InvalidPacketNumber, + firstCapablePacket: protocol.InvalidPacketNumber, + state: ecnStateInitial, + logger: logger, + qlogger: qlogger, + } +} + +func (e *ecnTracker) SentPacket(pn protocol.PacketNumber, ecn protocol.ECN) { + //nolint:exhaustive // These are the only ones we need to take care of. + switch ecn { + case protocol.ECNNon: + return + case protocol.ECT0: + e.numSentECT0++ + case protocol.ECT1: + e.numSentECT1++ + case protocol.ECNUnsupported: + if e.state != ecnStateFailed { + panic("didn't expect ECN to be unsupported") + } + default: + panic(fmt.Sprintf("sent packet with unexpected ECN marking: %s", ecn)) + } + + if e.state == ecnStateCapable && e.firstCapablePacket == protocol.InvalidPacketNumber { + e.firstCapablePacket = pn + } + + if e.state != ecnStateTesting { + return + } + + e.numSentTesting++ + if e.firstTestingPacket == protocol.InvalidPacketNumber { + e.firstTestingPacket = pn + } + if e.numSentECT0+e.numSentECT1 >= numECNTestingPackets { + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateUnknown, + }) + } + e.state = ecnStateUnknown + e.lastTestingPacket = pn + } +} + +func (e *ecnTracker) Mode() protocol.ECN { + switch e.state { + case ecnStateInitial: + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateTesting, + }) + } + e.state = ecnStateTesting + return e.Mode() + case ecnStateTesting, ecnStateCapable: + return protocol.ECT0 + case ecnStateUnknown, ecnStateFailed: + return protocol.ECNNon + default: + panic(fmt.Sprintf("unknown ECN state: %d", e.state)) + } +} + +func (e *ecnTracker) LostPacket(pn protocol.PacketNumber) { + if e.state != ecnStateTesting && e.state != ecnStateUnknown { + return + } + if !e.isTestingPacket(pn) { + return + } + e.numLostTesting++ + // Only proceed if we have sent all 10 testing packets. + if e.state != ecnStateUnknown { + return + } + if e.numLostTesting >= e.numSentTesting { + e.logger.Debugf("Disabling ECN. All testing packets were lost.") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedLostAllTestingPackets, + }) + } + e.state = ecnStateFailed + return + } + // Path validation also fails if some testing packets are lost, and all other testing packets where CE-marked + e.failIfMangled() +} + +// HandleNewlyAcked handles the ECN counts on an ACK frame. +// It must only be called for ACK frames that increase the largest acknowledged packet number, +// see section 13.4.2.1 of RFC 9000. +func (e *ecnTracker) HandleNewlyAcked(packets []packetWithPacketNumber, ect0, ect1, ecnce int64) (congested bool) { + if e.state == ecnStateFailed { + return false + } + + // ECN validation can fail if the received total count for either ECT(0) or ECT(1) exceeds + // the total number of packets sent with each corresponding ECT codepoint. + if ect0 > e.numSentECT0 || ect1 > e.numSentECT1 { + e.logger.Debugf("Disabling ECN. Received more ECT(0) / ECT(1) acknowledgements than packets sent.") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedMoreECNCountsThanSent, + }) + } + e.state = ecnStateFailed + return false + } + + // Count ECT0 and ECT1 marks that we used when sending the packets that are now being acknowledged. + var ackedECT0, ackedECT1 int64 + for _, p := range packets { + //nolint:exhaustive // We only ever send ECT(0) and ECT(1). + switch e.ecnMarking(p.PacketNumber) { + case protocol.ECT0: + ackedECT0++ + case protocol.ECT1: + ackedECT1++ + } + } + + // If an ACK frame newly acknowledges a packet that the endpoint sent with either the ECT(0) or ECT(1) + // codepoint set, ECN validation fails if the corresponding ECN counts are not present in the ACK frame. + // This check detects: + // * paths that bleach all ECN marks, and + // * peers that don't report any ECN counts + if (ackedECT0 > 0 || ackedECT1 > 0) && ect0 == 0 && ect1 == 0 && ecnce == 0 { + e.logger.Debugf("Disabling ECN. ECN-marked packet acknowledged, but no ECN counts on ACK frame.") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedNoECNCounts, + }) + } + e.state = ecnStateFailed + return false + } + + // Determine the increase in ECT0, ECT1 and ECNCE marks + newECT0 := ect0 - e.numAckedECT0 + newECT1 := ect1 - e.numAckedECT1 + newECNCE := ecnce - e.numAckedECNCE + + // We're only processing ACKs that increase the Largest Acked. + // Therefore, the ECN counters should only ever increase. + // Any decrease means that the peer's counting logic is broken. + if newECT0 < 0 || newECT1 < 0 || newECNCE < 0 { + e.logger.Debugf("Disabling ECN. ECN counts decreased unexpectedly.") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedDecreasedECNCounts, + }) + } + e.state = ecnStateFailed + return false + } + + // ECN validation also fails if the sum of the increase in ECT(0) and ECN-CE counts is less than the number + // of newly acknowledged packets that were originally sent with an ECT(0) marking. + // This could be the result of (partial) bleaching. + if newECT0+newECNCE < ackedECT0 { + e.logger.Debugf("Disabling ECN. Received less ECT(0) + ECN-CE than packets sent with ECT(0).") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedTooFewECNCounts, + }) + } + e.state = ecnStateFailed + return false + } + // Similarly, ECN validation fails if the sum of the increases to ECT(1) and ECN-CE counts is less than + // the number of newly acknowledged packets sent with an ECT(1) marking. + if newECT1+newECNCE < ackedECT1 { + e.logger.Debugf("Disabling ECN. Received less ECT(1) + ECN-CE than packets sent with ECT(1).") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedTooFewECNCounts, + }) + } + e.state = ecnStateFailed + return false + } + + // update our counters + e.numAckedECT0 = ect0 + e.numAckedECT1 = ect1 + e.numAckedECNCE = ecnce + + // Detect mangling (a path remarking all ECN-marked testing packets as CE), + // once all 10 testing packets have been sent out. + if e.state == ecnStateUnknown { + e.failIfMangled() + if e.state == ecnStateFailed { + return false + } + } + if e.state == ecnStateTesting || e.state == ecnStateUnknown { + var ackedTestingPacket bool + for _, p := range packets { + if e.isTestingPacket(p.PacketNumber) { + ackedTestingPacket = true + break + } + } + // This check won't succeed if the path is mangling ECN-marks (i.e. rewrites all ECN-marked packets to CE). + if ackedTestingPacket && (newECT0 > 0 || newECT1 > 0) { + e.logger.Debugf("ECN capability confirmed.") + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateCapable, + }) + } + e.state = ecnStateCapable + } + } + + // Don't trust CE marks before having confirmed ECN capability of the path. + // Otherwise, mangling would be misinterpreted as actual congestion. + return e.state == ecnStateCapable && newECNCE > 0 +} + +// failIfMangled fails ECN validation if all testing packets are lost or CE-marked. +func (e *ecnTracker) failIfMangled() { + numAckedECNCE := e.numAckedECNCE + int64(e.numLostTesting) + if e.numSentECT0+e.numSentECT1 > numAckedECNCE { + return + } + if e.qlogger != nil { + e.qlogger.RecordEvent(qlog.ECNStateUpdated{ + State: qlog.ECNStateFailed, + Trigger: ecnFailedManglingDetected, + }) + } + e.state = ecnStateFailed +} + +func (e *ecnTracker) ecnMarking(pn protocol.PacketNumber) protocol.ECN { + if pn < e.firstTestingPacket || e.firstTestingPacket == protocol.InvalidPacketNumber { + return protocol.ECNNon + } + if pn < e.lastTestingPacket || e.lastTestingPacket == protocol.InvalidPacketNumber { + return protocol.ECT0 + } + if pn < e.firstCapablePacket || e.firstCapablePacket == protocol.InvalidPacketNumber { + return protocol.ECNNon + } + // We don't need to deal with the case when ECN validation fails, + // since we're ignoring any ECN counts reported in ACK frames in that case. + return protocol.ECT0 +} + +func (e *ecnTracker) isTestingPacket(pn protocol.PacketNumber) bool { + if e.firstTestingPacket == protocol.InvalidPacketNumber { + return false + } + return pn >= e.firstTestingPacket && (pn <= e.lastTestingPacket || e.lastTestingPacket == protocol.InvalidPacketNumber) +} diff --git a/third_party/quic-go/internal/ackhandler/ecn_test.go b/third_party/quic-go/internal/ackhandler/ecn_test.go new file mode 100644 index 0000000..757a548 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/ecn_test.go @@ -0,0 +1,353 @@ +package ackhandler + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func getAckedPackets(pns ...protocol.PacketNumber) []packetWithPacketNumber { + var packets []packetWithPacketNumber + for _, p := range pns { + packets = append(packets, packetWithPacketNumber{PacketNumber: p}) + } + return packets +} + +// sendECNTestingPackets sends 10 ECT(0) packets, and then one more packet +// Packet numbers: 0 through 9. +func sendECNTestingPackets(t *testing.T, ecnTracker *ecnTracker, recorder *events.Recorder) { + t.Helper() + + for i := range protocol.PacketNumber(9) { + require.Equal(t, protocol.ECT0, ecnTracker.Mode()) + // do this twice to make sure only sent packets are counted + require.Equal(t, protocol.ECT0, ecnTracker.Mode()) + ecnTracker.SentPacket(i, protocol.ECT0) + } + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateTesting}}, + recorder.Events(), + ) + require.Equal(t, protocol.ECT0, ecnTracker.Mode()) + recorder.Clear() + ecnTracker.SentPacket(9, protocol.ECT0) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateUnknown}}, + recorder.Events(), + ) + recorder.Clear() + // in unknown state, packets shouldn't be ECN-marked + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) +} + +// ECN validation fails if *all* ECN testing packets are lost. +func TestECNTestingPacketsLoss(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + + // send non-testing packets + for i := range protocol.PacketNumber(10) { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(10+i, protocol.ECNNon) + } + + // lose all but one packet + for pn := range protocol.PacketNumber(10) { + if pn == 4 { + continue + } + ecnTracker.LostPacket(pn) + } + // loss of non-testing packets doesn't matter + ecnTracker.LostPacket(13) + ecnTracker.LostPacket(14) + + // now lose the last testing packet + require.Empty(t, eventRecorder.Events()) + eventRecorder.Clear() + ecnTracker.LostPacket(4) + require.Equal(t, + []qlogwriter.Event{ + qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: ecnFailedLostAllTestingPackets}, + }, + eventRecorder.Events(), + ) +} + +// ECN support is validated once an acknowledgment for any testing packet is received. +// This applies even if that happens before all testing packets have been sent out. +func TestECNValidationInTestingState(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + for i := range 5 { + require.Equal(t, protocol.ECT0, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECT0) + } + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateTesting}}, + eventRecorder.Events(), + ) + eventRecorder.Clear() + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(3), 1, 0, 0)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateCapable}}, + eventRecorder.Events(), + ) + + // make sure we continue sending ECT(0) packets + for i := 5; i < 100; i++ { + require.Equal(t, protocol.ECT0, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECT0) + } +} + +// ENC is also validated after all testing packets have been sent out, +// once an acknowledgment for any testing packet is received. +func TestECNValidationInUnknownState(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + + for i := range protocol.PacketNumber(10) { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + pn := 10 + i + ecnTracker.SentPacket(pn, protocol.ECNNon) + // lose some packets to make sure this doesn't influence the outcome. + if i%2 == 0 { + ecnTracker.LostPacket(pn) + } + } + require.Empty(t, eventRecorder.Events()) + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(7), 1, 0, 0)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateCapable}}, + eventRecorder.Events(), + ) +} + +func TestECNValidationFailures(t *testing.T) { + t.Run("ECN bleaching", func(t *testing.T) { + // this ACK doesn't contain any ECN counts + testECNValidationFailure(t, getAckedPackets(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12), 0, 0, 0, ecnFailedNoECNCounts) + }) + + t.Run("wrong ECN code point", func(t *testing.T) { + // we sent ECT(0), but this ACK acknowledges ECT(1) + testECNValidationFailure(t, getAckedPackets(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12), 0, 1, 0, ecnFailedMoreECNCountsThanSent) + }) + + t.Run("more ECN counts than sent packets", func(t *testing.T) { + // only 10 ECT(0) packets were sent, but the ACK claims to have received 12 of them + testECNValidationFailure(t, getAckedPackets(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12), 12, 0, 0, ecnFailedMoreECNCountsThanSent) + }) +} + +func testECNValidationFailure( + t *testing.T, + ackedPackets []packetWithPacketNumber, + ect0, ect1, ecnce int64, + expectedTrigger string, +) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + + require.False(t, ecnTracker.HandleNewlyAcked(ackedPackets, ect0, ect1, ecnce)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: expectedTrigger}}, + eventRecorder.Events(), + ) +} + +func TestECNValidationNotEnoughECNCounts(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + require.Empty(t, eventRecorder.Events()) + // First only acknowledge some packets sent with ECN marks. + require.True(t, ecnTracker.HandleNewlyAcked(getAckedPackets(1, 2, 3, 12), 2, 0, 1)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateCapable}}, + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // Now acknowledge some more packets sent with ECN marks, but don't increase the counters enough. + // This ACK acknowledges 3 more ECN-marked packets, but the counters only increase by 2. + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(4, 5, 6, 15), 3, 0, 2)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: ecnFailedTooFewECNCounts}}, + eventRecorder.Events(), + ) +} + +func TestECNNonsensicalECNCountDecrease(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + require.Empty(t, eventRecorder.Events()) + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(1, 2, 3, 12), 3, 0, 0)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateCapable}}, + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // Now acknowledge some more packets, but decrease the ECN counts. Obviously, this doesn't make any sense. + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(4, 5, 6, 13), 2, 0, 0)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: ecnFailedDecreasedECNCounts}}, + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // make sure that new ACKs are ignored + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(7, 8, 9, 14), 5, 0, 0)) + require.Empty(t, eventRecorder.Events()) +} + +func TestECNACKReordering(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + require.Empty(t, eventRecorder.Events()) + + // The ACK contains more ECN counts than it acknowledges packets. + // This can happen if ACKs are lost / reordered. + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(1, 2, 3, 12), 8, 0, 0)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateCapable}}, + eventRecorder.Events(), + ) +} + +// Mangling is detected if all testing packets are marked CE. +func TestECNManglingAllPacketsMarkedCE(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + + // ECN capability not confirmed yet, therefore CE marks are not regarded as congestion events + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(0, 1, 2, 3), 0, 0, 4)) + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(4, 5, 6, 10, 11, 12), 0, 0, 7)) + require.Empty(t, eventRecorder.Events()) + + // With the next ACK, all testing packets will now have been marked CE. + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(7, 8, 9, 13), 0, 0, 10)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: ecnFailedManglingDetected}}, + eventRecorder.Events(), + ) +} + +// Mangling is also detected if some testing packets are lost, and then others are marked CE. +func TestECNManglingSomePacketsLostSomeMarkedCE(t *testing.T) { + t.Run("packet loss first", func(t *testing.T) { + testECNManglingSomePacketsLostSomeMarkedCE(t, true) + }) + t.Run("CE marking first", func(t *testing.T) { + testECNManglingSomePacketsLostSomeMarkedCE(t, false) + }) +} + +func testECNManglingSomePacketsLostSomeMarkedCE(t *testing.T, packetLossFirst bool) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + // Lose a few packets. + if packetLossFirst { + ecnTracker.LostPacket(0) + ecnTracker.LostPacket(1) + ecnTracker.LostPacket(2) + } + // ECN capability not confirmed yet, therefore CE marks are not regarded as congestion events + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(3, 4, 5, 6, 7, 8), 0, 0, 6)) + require.Empty(t, eventRecorder.Events()) + // By CE-marking the last unacknowledged testing packets, we should detect the mangling. + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(9), 0, 0, 7)) + if packetLossFirst { + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: ecnFailedManglingDetected}}, + eventRecorder.Events(), + ) + } else { + require.Empty(t, eventRecorder.Events()) + } + + if !packetLossFirst { + ecnTracker.LostPacket(0) + ecnTracker.LostPacket(1) + ecnTracker.LostPacket(2) + + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateFailed, Trigger: ecnFailedManglingDetected}}, + eventRecorder.Events(), + ) + } +} + +func TestECNCongestionDetection(t *testing.T) { + var eventRecorder events.Recorder + ecnTracker := newECNTracker(utils.DefaultLogger, &eventRecorder) + + sendECNTestingPackets(t, ecnTracker, &eventRecorder) + for i := 10; i < 20; i++ { + require.Equal(t, protocol.ECNNon, ecnTracker.Mode()) + ecnTracker.SentPacket(protocol.PacketNumber(i), protocol.ECNNon) + } + // Receive one CE count. + require.True(t, ecnTracker.HandleNewlyAcked(getAckedPackets(1, 2, 3, 12), 2, 0, 1)) + require.Equal(t, + []qlogwriter.Event{qlog.ECNStateUpdated{State: qlog.ECNStateCapable}}, + eventRecorder.Events(), + ) + + // No increase in CE. No congestion. + require.False(t, ecnTracker.HandleNewlyAcked(getAckedPackets(4, 5, 6, 13), 5, 0, 1)) + eventRecorder.Clear() + + // Increase in CE. More congestion. + require.True(t, ecnTracker.HandleNewlyAcked(getAckedPackets(7, 8, 9, 14), 7, 0, 2)) + require.Empty(t, eventRecorder.Events()) +} diff --git a/third_party/quic-go/internal/ackhandler/frame.go b/third_party/quic-go/internal/ackhandler/frame.go new file mode 100644 index 0000000..edab8df --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/frame.go @@ -0,0 +1,21 @@ +package ackhandler + +import ( + "github.com/apernet/quic-go/internal/wire" +) + +// FrameHandler handles the acknowledgement and the loss of a frame. +type FrameHandler interface { + OnAcked(wire.Frame) + OnLost(wire.Frame) +} + +type Frame struct { + Frame wire.Frame // nil if the frame has already been acknowledged in another packet + Handler FrameHandler +} + +type StreamFrame struct { + Frame *wire.StreamFrame + Handler FrameHandler +} diff --git a/third_party/quic-go/internal/ackhandler/interfaces.go b/third_party/quic-go/internal/ackhandler/interfaces.go new file mode 100644 index 0000000..82046b8 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/interfaces.go @@ -0,0 +1,45 @@ +package ackhandler + +import ( + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +// SentPacketHandler handles ACKs received for outgoing packets +type SentPacketHandler interface { + // SentPacket may modify the packet + SentPacket(t monotime.Time, pn, largestAcked protocol.PacketNumber, streamFrames []StreamFrame, frames []Frame, encLevel protocol.EncryptionLevel, ecn protocol.ECN, size protocol.ByteCount, isPathMTUProbePacket, isPathProbePacket bool) + // ReceivedAck processes an ACK frame. + // It does not store a copy of the frame. + ReceivedAck(f *wire.AckFrame, encLevel protocol.EncryptionLevel, rcvTime monotime.Time) (bool /* 1-RTT packet acked */, error) + ReceivedPacket(protocol.EncryptionLevel, monotime.Time) + ReceivedBytes(_ protocol.ByteCount, rcvTime monotime.Time) + DropPackets(_ protocol.EncryptionLevel, rcvTime monotime.Time) + ResetForRetry(rcvTime monotime.Time) + + // The SendMode determines if and what kind of packets can be sent. + SendMode(now monotime.Time) SendMode + // TimeUntilSend is the time when the next packet should be sent. + // It is used for pacing packets. + TimeUntilSend() monotime.Time + SetMaxDatagramSize(count protocol.ByteCount) + // SetLastDatagramPadding reports how much room was left in the datagram that + // was just packed. It only matters when packet numbers are shortened. + SetLastDatagramPadding(protocol.ByteCount) + + // only to be called once the handshake is complete + QueueProbePacket(protocol.EncryptionLevel) bool /* was a packet queued */ + + ECNMode(isShortHeaderPacket bool) protocol.ECN // isShortHeaderPacket should only be true for non-coalesced 1-RTT packets + PeekPacketNumber(protocol.EncryptionLevel) (protocol.PacketNumber, protocol.PacketNumberLen) + PopPacketNumber(protocol.EncryptionLevel) protocol.PacketNumber + + GetLossDetectionTimeout() monotime.Time + OnLossDetectionTimeout(now monotime.Time) error + + MigratedPath(now monotime.Time, initialMaxPacketSize protocol.ByteCount) + + SetCongestionControl(congestion.CongestionControl) +} diff --git a/third_party/quic-go/internal/ackhandler/lost_packet_tracker.go b/third_party/quic-go/internal/ackhandler/lost_packet_tracker.go new file mode 100644 index 0000000..05de0d9 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/lost_packet_tracker.go @@ -0,0 +1,73 @@ +package ackhandler + +import ( + "iter" + "slices" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +type lostPacket struct { + PacketNumber protocol.PacketNumber + SendTime monotime.Time +} + +type lostPacketTracker struct { + maxLength int + lostPackets []lostPacket +} + +func newLostPacketTracker(maxLength int) *lostPacketTracker { + return &lostPacketTracker{ + maxLength: maxLength, + // Preallocate a small slice only. + // Hopefully we won't lose many packets. + lostPackets: make([]lostPacket, 0, 4), + } +} + +func (t *lostPacketTracker) Add(p protocol.PacketNumber, sendTime monotime.Time) { + if len(t.lostPackets) == t.maxLength { + t.lostPackets = t.lostPackets[1:] + } + t.lostPackets = append(t.lostPackets, lostPacket{ + PacketNumber: p, + SendTime: sendTime, + }) +} + +// Delete deletes a packet from the lost packet tracker. +// This function is not optimized for performance if many packets are lost, +// but it is only used when a spurious loss is detected, which is rare. +func (t *lostPacketTracker) Delete(pn protocol.PacketNumber) { + t.lostPackets = slices.DeleteFunc(t.lostPackets, func(p lostPacket) bool { + return p.PacketNumber == pn + }) +} + +func (t *lostPacketTracker) All() iter.Seq2[protocol.PacketNumber, monotime.Time] { + return func(yield func(protocol.PacketNumber, monotime.Time) bool) { + for _, p := range t.lostPackets { + if !yield(p.PacketNumber, p.SendTime) { + return + } + } + } +} + +func (t *lostPacketTracker) DeleteBefore(ti monotime.Time) { + if len(t.lostPackets) == 0 { + return + } + if !t.lostPackets[0].SendTime.Before(ti) { + return + } + var idx int + for ; idx < len(t.lostPackets); idx++ { + if !t.lostPackets[idx].SendTime.Before(ti) { + break + } + } + t.lostPackets = slices.Delete(t.lostPackets, 0, idx) +} diff --git a/third_party/quic-go/internal/ackhandler/lost_packet_tracker_test.go b/third_party/quic-go/internal/ackhandler/lost_packet_tracker_test.go new file mode 100644 index 0000000..9dd6eba --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/lost_packet_tracker_test.go @@ -0,0 +1,75 @@ +package ackhandler + +import ( + "maps" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestLostPacketTracker(t *testing.T) { + lt := newLostPacketTracker(4) + + start := monotime.Now() + lt.Add(1, start) + lt.Add(5, start.Add(time.Second)) + lt.Add(8, start.Add(2*time.Second)) + require.Equal(t, map[protocol.PacketNumber]monotime.Time{ + 1: start, + 5: start.Add(time.Second), + 8: start.Add(2 * time.Second), + }, maps.Collect(lt.All())) + + // Lose 2 more packets. The first one should be removed. + lt.Add(10, start.Add(3*time.Second)) + lt.Add(11, start.Add(4*time.Second)) + require.Equal(t, map[protocol.PacketNumber]monotime.Time{ + 5: start.Add(time.Second), + 8: start.Add(2 * time.Second), + 10: start.Add(3 * time.Second), + 11: start.Add(4 * time.Second), + }, maps.Collect(lt.All())) + + lt.Delete(5) + lt.Delete(10) + require.Equal(t, map[protocol.PacketNumber]monotime.Time{ + 8: start.Add(2 * time.Second), + 11: start.Add(4 * time.Second), + }, maps.Collect(lt.All())) +} + +func TestLostPacketTrackerDeleteBefore(t *testing.T) { + lt := newLostPacketTracker(4) + + trackedPackets := func(lt *lostPacketTracker) []protocol.PacketNumber { + var pns []protocol.PacketNumber + for pn := range lt.All() { + pns = append(pns, pn) + } + return pns + } + + start := monotime.Now() + lt.Add(1, start) + lt.Add(5, start.Add(time.Second)) + lt.Add(8, start.Add(2*time.Second)) + lt.Add(10, start.Add(3*time.Second)) + + require.Equal(t, []protocol.PacketNumber{1, 5, 8, 10}, trackedPackets(lt)) + + lt.DeleteBefore(start) // this should be a no-op + require.Equal(t, []protocol.PacketNumber{1, 5, 8, 10}, trackedPackets(lt)) + + lt.DeleteBefore(start.Add(2 * time.Second)) + require.Equal(t, []protocol.PacketNumber{8, 10}, trackedPackets(lt)) + + lt.DeleteBefore(start.Add(time.Second * 5 / 2)) + require.Equal(t, []protocol.PacketNumber{10}, trackedPackets(lt)) + + lt.DeleteBefore(start.Add(time.Hour)) + require.Empty(t, trackedPackets(lt)) +} diff --git a/third_party/quic-go/internal/ackhandler/mock_ecn_handler_test.go b/third_party/quic-go/internal/ackhandler/mock_ecn_handler_test.go new file mode 100644 index 0000000..89469d7 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/mock_ecn_handler_test.go @@ -0,0 +1,189 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/ackhandler (interfaces: ECNHandler) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package ackhandler -destination mock_ecn_handler_test.go github.com/apernet/quic-go/internal/ackhandler ECNHandler +// + +// Package ackhandler is a generated GoMock package. +package ackhandler + +import ( + reflect "reflect" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockECNHandler is a mock of ECNHandler interface. +type MockECNHandler struct { + ctrl *gomock.Controller + recorder *MockECNHandlerMockRecorder + isgomock struct{} +} + +// MockECNHandlerMockRecorder is the mock recorder for MockECNHandler. +type MockECNHandlerMockRecorder struct { + mock *MockECNHandler +} + +// NewMockECNHandler creates a new mock instance. +func NewMockECNHandler(ctrl *gomock.Controller) *MockECNHandler { + mock := &MockECNHandler{ctrl: ctrl} + mock.recorder = &MockECNHandlerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockECNHandler) EXPECT() *MockECNHandlerMockRecorder { + return m.recorder +} + +// HandleNewlyAcked mocks base method. +func (m *MockECNHandler) HandleNewlyAcked(packets []packetWithPacketNumber, ect0, ect1, ecnce int64) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "HandleNewlyAcked", packets, ect0, ect1, ecnce) + ret0, _ := ret[0].(bool) + return ret0 +} + +// HandleNewlyAcked indicates an expected call of HandleNewlyAcked. +func (mr *MockECNHandlerMockRecorder) HandleNewlyAcked(packets, ect0, ect1, ecnce any) *MockECNHandlerHandleNewlyAckedCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HandleNewlyAcked", reflect.TypeOf((*MockECNHandler)(nil).HandleNewlyAcked), packets, ect0, ect1, ecnce) + return &MockECNHandlerHandleNewlyAckedCall{Call: call} +} + +// MockECNHandlerHandleNewlyAckedCall wrap *gomock.Call +type MockECNHandlerHandleNewlyAckedCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockECNHandlerHandleNewlyAckedCall) Return(congested bool) *MockECNHandlerHandleNewlyAckedCall { + c.Call = c.Call.Return(congested) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockECNHandlerHandleNewlyAckedCall) Do(f func([]packetWithPacketNumber, int64, int64, int64) bool) *MockECNHandlerHandleNewlyAckedCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockECNHandlerHandleNewlyAckedCall) DoAndReturn(f func([]packetWithPacketNumber, int64, int64, int64) bool) *MockECNHandlerHandleNewlyAckedCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// LostPacket mocks base method. +func (m *MockECNHandler) LostPacket(arg0 protocol.PacketNumber) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "LostPacket", arg0) +} + +// LostPacket indicates an expected call of LostPacket. +func (mr *MockECNHandlerMockRecorder) LostPacket(arg0 any) *MockECNHandlerLostPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LostPacket", reflect.TypeOf((*MockECNHandler)(nil).LostPacket), arg0) + return &MockECNHandlerLostPacketCall{Call: call} +} + +// MockECNHandlerLostPacketCall wrap *gomock.Call +type MockECNHandlerLostPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockECNHandlerLostPacketCall) Return() *MockECNHandlerLostPacketCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockECNHandlerLostPacketCall) Do(f func(protocol.PacketNumber)) *MockECNHandlerLostPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockECNHandlerLostPacketCall) DoAndReturn(f func(protocol.PacketNumber)) *MockECNHandlerLostPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Mode mocks base method. +func (m *MockECNHandler) Mode() protocol.ECN { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Mode") + ret0, _ := ret[0].(protocol.ECN) + return ret0 +} + +// Mode indicates an expected call of Mode. +func (mr *MockECNHandlerMockRecorder) Mode() *MockECNHandlerModeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Mode", reflect.TypeOf((*MockECNHandler)(nil).Mode)) + return &MockECNHandlerModeCall{Call: call} +} + +// MockECNHandlerModeCall wrap *gomock.Call +type MockECNHandlerModeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockECNHandlerModeCall) Return(arg0 protocol.ECN) *MockECNHandlerModeCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockECNHandlerModeCall) Do(f func() protocol.ECN) *MockECNHandlerModeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockECNHandlerModeCall) DoAndReturn(f func() protocol.ECN) *MockECNHandlerModeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SentPacket mocks base method. +func (m *MockECNHandler) SentPacket(arg0 protocol.PacketNumber, arg1 protocol.ECN) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SentPacket", arg0, arg1) +} + +// SentPacket indicates an expected call of SentPacket. +func (mr *MockECNHandlerMockRecorder) SentPacket(arg0, arg1 any) *MockECNHandlerSentPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SentPacket", reflect.TypeOf((*MockECNHandler)(nil).SentPacket), arg0, arg1) + return &MockECNHandlerSentPacketCall{Call: call} +} + +// MockECNHandlerSentPacketCall wrap *gomock.Call +type MockECNHandlerSentPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockECNHandlerSentPacketCall) Return() *MockECNHandlerSentPacketCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockECNHandlerSentPacketCall) Do(f func(protocol.PacketNumber, protocol.ECN)) *MockECNHandlerSentPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockECNHandlerSentPacketCall) DoAndReturn(f func(protocol.PacketNumber, protocol.ECN)) *MockECNHandlerSentPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/ackhandler/mockgen.go b/third_party/quic-go/internal/ackhandler/mockgen.go new file mode 100644 index 0000000..144563c --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/mockgen.go @@ -0,0 +1,6 @@ +//go:build gomock || generate + +package ackhandler + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package ackhandler -destination mock_ecn_handler_test.go github.com/apernet/quic-go/internal/ackhandler ECNHandler" +type ECNHandler = ecnHandler diff --git a/third_party/quic-go/internal/ackhandler/packet.go b/third_party/quic-go/internal/ackhandler/packet.go new file mode 100644 index 0000000..5b57a14 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/packet.go @@ -0,0 +1,60 @@ +package ackhandler + +import ( + "sync" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +type packetWithPacketNumber struct { + PacketNumber protocol.PacketNumber + *packet +} + +// A Packet is a packet +type packet struct { + SendTime monotime.Time + StreamFrames []StreamFrame + Frames []Frame + LargestAcked protocol.PacketNumber // InvalidPacketNumber if the packet doesn't contain an ACK + Length protocol.ByteCount + EncryptionLevel protocol.EncryptionLevel + + IsPathMTUProbePacket bool // We don't report the loss of Path MTU probe packets to the congestion controller. + + includedInBytesInFlight bool + isPathProbePacket bool +} + +func (p *packet) Outstanding() bool { + return !p.IsPathMTUProbePacket && !p.isPathProbePacket && p.IsAckEliciting() +} + +func (p *packet) IsAckEliciting() bool { + return len(p.StreamFrames) > 0 || len(p.Frames) > 0 +} + +var packetPool = sync.Pool{New: func() any { return &packet{} }} + +func getPacket() *packet { + p := packetPool.Get().(*packet) + p.StreamFrames = nil + p.Frames = nil + p.LargestAcked = 0 + p.Length = 0 + p.EncryptionLevel = protocol.EncryptionLevel(0) + p.SendTime = 0 + p.IsPathMTUProbePacket = false + p.includedInBytesInFlight = false + p.isPathProbePacket = false + return p +} + +// We currently only return Packets back into the pool when they're acknowledged (not when they're lost). +// This simplifies the code, and gives the vast majority of the performance benefit we can gain from using the pool. +func putPacket(p *packet) { + p.Frames = nil + p.StreamFrames = nil + packetPool.Put(p) +} diff --git a/third_party/quic-go/internal/ackhandler/packet_number_generator.go b/third_party/quic-go/internal/ackhandler/packet_number_generator.go new file mode 100644 index 0000000..2e5fac3 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/packet_number_generator.go @@ -0,0 +1,84 @@ +package ackhandler + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +type packetNumberGenerator interface { + Peek() protocol.PacketNumber + // Pop pops the packet number. + // It reports if the packet number (before the one just popped) was skipped. + // It never skips more than one packet number in a row. + Pop() (skipped bool, _ protocol.PacketNumber) +} + +type sequentialPacketNumberGenerator struct { + next protocol.PacketNumber +} + +var _ packetNumberGenerator = &sequentialPacketNumberGenerator{} + +func newSequentialPacketNumberGenerator(initial protocol.PacketNumber) packetNumberGenerator { + return &sequentialPacketNumberGenerator{next: initial} +} + +func (p *sequentialPacketNumberGenerator) Peek() protocol.PacketNumber { + return p.next +} + +func (p *sequentialPacketNumberGenerator) Pop() (bool, protocol.PacketNumber) { + next := p.next + p.next++ + return false, next +} + +// The skippingPacketNumberGenerator generates the packet number for the next packet +// it randomly skips a packet number every averagePeriod packets (on average). +// It is guaranteed to never skip two consecutive packet numbers. +type skippingPacketNumberGenerator struct { + period protocol.PacketNumber + maxPeriod protocol.PacketNumber + + next protocol.PacketNumber + nextToSkip protocol.PacketNumber + + rng utils.Rand +} + +var _ packetNumberGenerator = &skippingPacketNumberGenerator{} + +func newSkippingPacketNumberGenerator(initial, initialPeriod, maxPeriod protocol.PacketNumber) packetNumberGenerator { + g := &skippingPacketNumberGenerator{ + next: initial, + period: initialPeriod, + maxPeriod: maxPeriod, + } + g.generateNewSkip() + return g +} + +func (p *skippingPacketNumberGenerator) Peek() protocol.PacketNumber { + if p.next == p.nextToSkip { + return p.next + 1 + } + return p.next +} + +func (p *skippingPacketNumberGenerator) Pop() (bool, protocol.PacketNumber) { + next := p.next + if p.next == p.nextToSkip { + next++ + p.next += 2 + p.generateNewSkip() + return true, next + } + p.next++ // generate a new packet number for the next packet + return false, next +} + +func (p *skippingPacketNumberGenerator) generateNewSkip() { + // make sure that there are never two consecutive packet numbers that are skipped + p.nextToSkip = p.next + 3 + protocol.PacketNumber(p.rng.Int31n(int32(2*p.period))) + p.period = min(2*p.period, p.maxPeriod) +} diff --git a/third_party/quic-go/internal/ackhandler/packet_number_generator_test.go b/third_party/quic-go/internal/ackhandler/packet_number_generator_test.go new file mode 100644 index 0000000..d9e7d7d --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/packet_number_generator_test.go @@ -0,0 +1,92 @@ +package ackhandler + +import ( + "math" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestSequentialPacketNumberGenerator(t *testing.T) { + const initialPN protocol.PacketNumber = 123 + png := newSequentialPacketNumberGenerator(initialPN) + + for i := initialPN; i < initialPN+1000; i++ { + require.Equal(t, i, png.Peek()) + require.Equal(t, i, png.Peek()) + skipNext, pn := png.Pop() + require.False(t, skipNext) + require.Equal(t, i, pn) + } +} + +func TestSkippingPacketNumberGenerator(t *testing.T) { + // the maximum period must be sufficiently small such that using a 32-bit random number is ok + require.Less(t, 2*protocol.SkipPacketMaxPeriod, protocol.PacketNumber(math.MaxInt32)) + + const initialPeriod protocol.PacketNumber = 25 + const maxPeriod protocol.PacketNumber = 300 + + png := newSkippingPacketNumberGenerator(100, initialPeriod, maxPeriod) + require.Equal(t, protocol.PacketNumber(100), png.Peek()) + require.Equal(t, protocol.PacketNumber(100), png.Peek()) + require.Equal(t, protocol.PacketNumber(100), png.Peek()) + _, pn := png.Pop() + require.Equal(t, protocol.PacketNumber(100), pn) + + var last protocol.PacketNumber + var skipped bool + for i := range maxPeriod { + didSkip, num := png.Pop() + if didSkip { + skipped = true + _, nextNum := png.Pop() + require.Equal(t, num+1, nextNum) + break + } + if i != 0 { + require.Equal(t, num, last+1) + } + last = num + } + require.True(t, skipped) +} + +func TestSkippingPacketNumberGeneratorPeriods(t *testing.T) { + const initialPN protocol.PacketNumber = 8 + const initialPeriod protocol.PacketNumber = 25 + const maxPeriod protocol.PacketNumber = 300 + + const rep = 2500 + periods := make([][]protocol.PacketNumber, rep) + expectedPeriods := []protocol.PacketNumber{25, 50, 100, 200, 300, 300, 300} + + for i := range rep { + png := newSkippingPacketNumberGenerator(initialPN, initialPeriod, maxPeriod) + lastSkip := initialPN + for len(periods[i]) < len(expectedPeriods) { + skipNext, next := png.Pop() + if skipNext { + skipped := next + 1 + require.Greater(t, skipped, lastSkip+1) + periods[i] = append(periods[i], skipped-lastSkip-1) + lastSkip = skipped + } + } + } + + for j := range expectedPeriods { + var average float64 + for i := range rep { + average += float64(periods[i][j]) / float64(len(periods)) + } + t.Logf("Period %d: %.2f (expected %d)\n", j, average, expectedPeriods[j]) + require.InDelta(t, + float64(expectedPeriods[j]+1), + average, + float64(max(protocol.PacketNumber(5), expectedPeriods[j]/10)), + ) + } +} diff --git a/third_party/quic-go/internal/ackhandler/received_packet_handler.go b/third_party/quic-go/internal/ackhandler/received_packet_handler.go new file mode 100644 index 0000000..57ebb5a --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/received_packet_handler.go @@ -0,0 +1,119 @@ +package ackhandler + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" +) + +type ReceivedPacketHandler struct { + initialPackets *receivedPacketTracker + handshakePackets *receivedPacketTracker + appDataPackets appDataReceivedPacketTracker + + lowest1RTTPacket protocol.PacketNumber +} + +func NewReceivedPacketHandler(logger utils.Logger) *ReceivedPacketHandler { + return &ReceivedPacketHandler{ + initialPackets: newReceivedPacketTracker(), + handshakePackets: newReceivedPacketTracker(), + appDataPackets: *newAppDataReceivedPacketTracker(logger), + lowest1RTTPacket: protocol.InvalidPacketNumber, + } +} + +func (h *ReceivedPacketHandler) ReceivedPacket( + pn protocol.PacketNumber, + ecn protocol.ECN, + encLevel protocol.EncryptionLevel, + rcvTime monotime.Time, + ackEliciting bool, +) error { + switch encLevel { + case protocol.EncryptionInitial: + return h.initialPackets.ReceivedPacket(pn, ecn, ackEliciting) + case protocol.EncryptionHandshake: + // The Handshake packet number space might already have been dropped as a result + // of processing the CRYPTO frame that was contained in this packet. + if h.handshakePackets == nil { + return nil + } + return h.handshakePackets.ReceivedPacket(pn, ecn, ackEliciting) + case protocol.Encryption0RTT: + if h.lowest1RTTPacket != protocol.InvalidPacketNumber && pn > h.lowest1RTTPacket { + return fmt.Errorf("received packet number %d on a 0-RTT packet after receiving %d on a 1-RTT packet", pn, h.lowest1RTTPacket) + } + return h.appDataPackets.ReceivedPacket(pn, ecn, rcvTime, ackEliciting) + case protocol.Encryption1RTT: + if h.lowest1RTTPacket == protocol.InvalidPacketNumber || pn < h.lowest1RTTPacket { + h.lowest1RTTPacket = pn + } + return h.appDataPackets.ReceivedPacket(pn, ecn, rcvTime, ackEliciting) + default: + panic(fmt.Sprintf("received packet with unknown encryption level: %s", encLevel)) + } +} + +func (h *ReceivedPacketHandler) IgnorePacketsBelow(pn protocol.PacketNumber) { + h.appDataPackets.IgnoreBelow(pn) +} + +func (h *ReceivedPacketHandler) DropPackets(encLevel protocol.EncryptionLevel) { + //nolint:exhaustive // 1-RTT packet number space is never dropped. + switch encLevel { + case protocol.EncryptionInitial: + h.initialPackets = nil + case protocol.EncryptionHandshake: + h.handshakePackets = nil + case protocol.Encryption0RTT: + // Nothing to do here. + // If we are rejecting 0-RTT, no 0-RTT packets will have been decrypted. + default: + panic(fmt.Sprintf("Cannot drop keys for encryption level %s", encLevel)) + } +} + +func (h *ReceivedPacketHandler) GetAlarmTimeout() monotime.Time { + return h.appDataPackets.GetAlarmTimeout() +} + +func (h *ReceivedPacketHandler) GetAckFrame(encLevel protocol.EncryptionLevel, now monotime.Time, onlyIfQueued bool) *wire.AckFrame { + //nolint:exhaustive // 0-RTT packets can't contain ACK frames. + switch encLevel { + case protocol.EncryptionInitial: + if h.initialPackets != nil { + return h.initialPackets.GetAckFrame() + } + return nil + case protocol.EncryptionHandshake: + if h.handshakePackets != nil { + return h.handshakePackets.GetAckFrame() + } + return nil + case protocol.Encryption1RTT: + return h.appDataPackets.GetAckFrame(now, onlyIfQueued) + default: + // 0-RTT packets can't contain ACK frames + return nil + } +} + +func (h *ReceivedPacketHandler) IsPotentiallyDuplicate(pn protocol.PacketNumber, encLevel protocol.EncryptionLevel) bool { + switch encLevel { + case protocol.EncryptionInitial: + if h.initialPackets != nil { + return h.initialPackets.IsPotentiallyDuplicate(pn) + } + case protocol.EncryptionHandshake: + if h.handshakePackets != nil { + return h.handshakePackets.IsPotentiallyDuplicate(pn) + } + case protocol.Encryption0RTT, protocol.Encryption1RTT: + return h.appDataPackets.IsPotentiallyDuplicate(pn) + } + panic("unexpected encryption level") +} diff --git a/third_party/quic-go/internal/ackhandler/received_packet_handler_test.go b/third_party/quic-go/internal/ackhandler/received_packet_handler_test.go new file mode 100644 index 0000000..2d74694 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/received_packet_handler_test.go @@ -0,0 +1,144 @@ +package ackhandler + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestGenerateACKsForPacketNumberSpaces(t *testing.T) { + handler := NewReceivedPacketHandler(utils.DefaultLogger) + + now := monotime.Now() + sendTime := now.Add(-time.Second) + + require.NoError(t, handler.ReceivedPacket(2, protocol.ECT0, protocol.EncryptionInitial, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(1, protocol.ECT1, protocol.EncryptionHandshake, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(5, protocol.ECNCE, protocol.Encryption1RTT, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(3, protocol.ECT0, protocol.EncryptionInitial, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(2, protocol.ECT1, protocol.EncryptionHandshake, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(4, protocol.ECNCE, protocol.Encryption1RTT, sendTime, true)) + + // Initial + initialAck := handler.GetAckFrame(protocol.EncryptionInitial, now, true) + require.NotNil(t, initialAck) + require.Equal(t, []wire.AckRange{{Smallest: 2, Largest: 3}}, initialAck.AckRanges) + require.Zero(t, initialAck.DelayTime) + require.EqualValues(t, 2, initialAck.ECT0) + require.Zero(t, initialAck.ECT1) + require.Zero(t, initialAck.ECNCE) + + // Handshake + handshakeAck := handler.GetAckFrame(protocol.EncryptionHandshake, now, true) + require.NotNil(t, handshakeAck) + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 2}}, handshakeAck.AckRanges) + require.Zero(t, handshakeAck.DelayTime) + require.Zero(t, handshakeAck.ECT0) + require.EqualValues(t, 2, handshakeAck.ECT1) + require.Zero(t, handshakeAck.ECNCE) + + // 1-RTT + oneRTTAck := handler.GetAckFrame(protocol.Encryption1RTT, now, true) + require.NotNil(t, oneRTTAck) + require.Equal(t, []wire.AckRange{{Smallest: 4, Largest: 5}}, oneRTTAck.AckRanges) + require.Equal(t, time.Second, oneRTTAck.DelayTime) + require.Zero(t, oneRTTAck.ECT0) + require.Zero(t, oneRTTAck.ECT1) + require.EqualValues(t, 2, oneRTTAck.ECNCE) +} + +func TestReceive0RTTAnd1RTT(t *testing.T) { + handler := NewReceivedPacketHandler(utils.DefaultLogger) + + sendTime := monotime.Now().Add(-time.Second) + + require.NoError(t, handler.ReceivedPacket(2, protocol.ECNNon, protocol.Encryption0RTT, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(3, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + + ack := handler.GetAckFrame(protocol.Encryption1RTT, monotime.Now(), true) + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 2, Largest: 3}}, ack.AckRanges) + + // 0-RTT packets with higher packet numbers than 1-RTT packets are rejected... + require.Error(t, handler.ReceivedPacket(4, protocol.ECNNon, protocol.Encryption0RTT, sendTime, true)) + // ... but reordered 0-RTT packets are allowed + require.NoError(t, handler.ReceivedPacket(1, protocol.ECNNon, protocol.Encryption0RTT, sendTime, true)) +} + +func TestDropPackets(t *testing.T) { + handler := NewReceivedPacketHandler(utils.DefaultLogger) + + sendTime := monotime.Now().Add(-time.Second) + + require.NoError(t, handler.ReceivedPacket(2, protocol.ECNNon, protocol.EncryptionInitial, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(1, protocol.ECNNon, protocol.EncryptionHandshake, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(2, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + + // Initial + require.NotNil(t, handler.GetAckFrame(protocol.EncryptionInitial, monotime.Now(), true)) + handler.DropPackets(protocol.EncryptionInitial) + require.Nil(t, handler.GetAckFrame(protocol.EncryptionInitial, monotime.Now(), true)) + + // Handshake + require.NotNil(t, handler.GetAckFrame(protocol.EncryptionHandshake, monotime.Now(), true)) + handler.DropPackets(protocol.EncryptionHandshake) + require.Nil(t, handler.GetAckFrame(protocol.EncryptionHandshake, monotime.Now(), true)) + + // 1-RTT + require.NotNil(t, handler.GetAckFrame(protocol.Encryption1RTT, monotime.Now(), true)) + + // 0-RTT is a no-op + handler.DropPackets(protocol.Encryption0RTT) +} + +func TestAckRangePruning(t *testing.T) { + handler := NewReceivedPacketHandler(utils.DefaultLogger) + + sendTime := monotime.Now() + require.NoError(t, handler.ReceivedPacket(1, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + require.NoError(t, handler.ReceivedPacket(2, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + + ack := handler.GetAckFrame(protocol.Encryption1RTT, monotime.Now(), true) + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 2}}, ack.AckRanges) + + require.NoError(t, handler.ReceivedPacket(3, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + handler.IgnorePacketsBelow(2) + require.NoError(t, handler.ReceivedPacket(4, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + + ack = handler.GetAckFrame(protocol.Encryption1RTT, monotime.Now(), true) + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 2, Largest: 4}}, ack.AckRanges) +} + +func TestPacketDuplicateDetection(t *testing.T) { + handler := NewReceivedPacketHandler(utils.DefaultLogger) + sendTime := monotime.Now() + + // 1-RTT is tested separately at the end + encLevels := []protocol.EncryptionLevel{ + protocol.EncryptionInitial, + protocol.EncryptionHandshake, + protocol.Encryption0RTT, + } + + for _, encLevel := range encLevels { + // first, packet 3 is not a duplicate + require.False(t, handler.IsPotentiallyDuplicate(3, encLevel)) + require.NoError(t, handler.ReceivedPacket(3, protocol.ECNNon, encLevel, sendTime, true)) + // now packet 3 is considered a duplicate + require.True(t, handler.IsPotentiallyDuplicate(3, encLevel)) + } + + // 1-RTT + require.True(t, handler.IsPotentiallyDuplicate(3, protocol.Encryption1RTT)) + require.False(t, handler.IsPotentiallyDuplicate(4, protocol.Encryption1RTT)) + require.NoError(t, handler.ReceivedPacket(4, protocol.ECNNon, protocol.Encryption1RTT, sendTime, true)) + require.True(t, handler.IsPotentiallyDuplicate(4, protocol.Encryption1RTT)) +} diff --git a/third_party/quic-go/internal/ackhandler/received_packet_history.go b/third_party/quic-go/internal/ackhandler/received_packet_history.go new file mode 100644 index 0000000..2b23790 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/received_packet_history.go @@ -0,0 +1,159 @@ +package ackhandler + +import ( + "iter" + "slices" + + "github.com/apernet/quic-go/internal/protocol" +) + +// interval is an interval from one PacketNumber to the other +type interval struct { + Start protocol.PacketNumber + End protocol.PacketNumber +} + +// The receivedPacketHistory stores if a packet number has already been received. +// It generates ACK ranges which can be used to assemble an ACK frame. +// It does not store packet contents. +type receivedPacketHistory struct { + ranges []interval // maximum length: protocol.MaxNumAckRanges + + deletedBelow protocol.PacketNumber +} + +func newReceivedPacketHistory() *receivedPacketHistory { + return &receivedPacketHistory{ + deletedBelow: protocol.InvalidPacketNumber, + } +} + +// ReceivedPacket registers a packet with PacketNumber p and updates the ranges +func (h *receivedPacketHistory) ReceivedPacket(p protocol.PacketNumber) bool /* is a new packet (and not a duplicate / delayed packet) */ { + // ignore delayed packets, if we already deleted the range + if p < h.deletedBelow { + return false + } + + isNew := h.addToRanges(p) + // Delete old ranges, if we're tracking too many of them. + // This is a DoS defense against a peer that sends us too many gaps. + if len(h.ranges) > protocol.MaxNumAckRanges { + h.ranges = slices.Delete(h.ranges, 0, len(h.ranges)-protocol.MaxNumAckRanges) + } + return isNew +} + +func (h *receivedPacketHistory) addToRanges(p protocol.PacketNumber) bool /* is a new packet (and not a duplicate / delayed packet) */ { + if len(h.ranges) == 0 { + h.ranges = append(h.ranges, interval{Start: p, End: p}) + return true + } + + for i := len(h.ranges) - 1; i >= 0; i-- { + // p already included in an existing range. Nothing to do here + if p >= h.ranges[i].Start && p <= h.ranges[i].End { + return false + } + + if h.ranges[i].End == p-1 { // extend a range at the end + h.ranges[i].End = p + return true + } + if h.ranges[i].Start == p+1 { // extend a range at the beginning + h.ranges[i].Start = p + + if i > 0 && h.ranges[i-1].End+1 == h.ranges[i].Start { // merge two ranges + h.ranges[i-1].End = h.ranges[i].End + h.ranges = slices.Delete(h.ranges, i, i+1) + } + return true + } + + // create a new range after the current one + if p > h.ranges[i].End { + h.ranges = slices.Insert(h.ranges, i+1, interval{Start: p, End: p}) + return true + } + } + + // create a new range at the beginning + h.ranges = slices.Insert(h.ranges, 0, interval{Start: p, End: p}) + return true +} + +// DeleteBelow deletes all entries below (but not including) p +func (h *receivedPacketHistory) DeleteBelow(p protocol.PacketNumber) { + if p < h.deletedBelow { + return + } + h.deletedBelow = p + + if len(h.ranges) == 0 { + return + } + + idx := -1 + for i := 0; i < len(h.ranges); i++ { + if h.ranges[i].End < p { // delete a whole range + idx = i + } else if p > h.ranges[i].Start && p <= h.ranges[i].End { + h.ranges[i].Start = p + break + } else { // no ranges affected. Nothing to do + break + } + } + if idx >= 0 { + h.ranges = slices.Delete(h.ranges, 0, idx+1) + } +} + +// Backward returns an iterator over the ranges in reverse order +func (h *receivedPacketHistory) Backward() iter.Seq[interval] { + return func(yield func(interval) bool) { + for i := len(h.ranges) - 1; i >= 0; i-- { + if !yield(h.ranges[i]) { + return + } + } + } +} + +func (h *receivedPacketHistory) HighestMissingUpTo(p protocol.PacketNumber) protocol.PacketNumber { + if len(h.ranges) == 0 || (h.deletedBelow != protocol.InvalidPacketNumber && p < h.deletedBelow) { + return protocol.InvalidPacketNumber + } + p = min(h.ranges[len(h.ranges)-1].End, p) + for i := len(h.ranges) - 1; i >= 0; i-- { + r := h.ranges[i] + if p >= r.Start && p <= r.End { // p is contained in this range + highest := r.Start - 1 // highest packet in the gap before this range + if h.deletedBelow != protocol.InvalidPacketNumber && highest < h.deletedBelow { + return protocol.InvalidPacketNumber + } + return highest + } + if i >= 1 && p > h.ranges[i-1].End && p <= r.Start { + // p is in the gap between the previous range and this range + return p + } + } + return p +} + +func (h *receivedPacketHistory) IsPotentiallyDuplicate(p protocol.PacketNumber) bool { + if p < h.deletedBelow { + return true + } + // Iterating over the slices is faster than using a binary search (using slices.BinarySearchFunc). + for i := len(h.ranges) - 1; i >= 0; i-- { + if p > h.ranges[i].End { + return false + } + if p <= h.ranges[i].End && p >= h.ranges[i].Start { + return true + } + } + return false +} diff --git a/third_party/quic-go/internal/ackhandler/received_packet_history_test.go b/third_party/quic-go/internal/ackhandler/received_packet_history_test.go new file mode 100644 index 0000000..03329fc --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/received_packet_history_test.go @@ -0,0 +1,304 @@ +package ackhandler + +import ( + "math/rand/v2" + "slices" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestReceivedPacketHistorySingleRange(t *testing.T) { + hist := newReceivedPacketHistory() + + require.True(t, hist.ReceivedPacket(4)) + require.Equal(t, []interval{{Start: 4, End: 4}}, slices.Collect(hist.Backward())) + + // add a duplicate packet + require.False(t, hist.ReceivedPacket(4)) + require.Equal(t, []interval{{Start: 4, End: 4}}, slices.Collect(hist.Backward())) + + // add a few more packets to extend the range + require.True(t, hist.ReceivedPacket(5)) + require.True(t, hist.ReceivedPacket(6)) + require.Equal(t, []interval{{Start: 4, End: 6}}, slices.Collect(hist.Backward())) + + // add a duplicate within this range + require.False(t, hist.ReceivedPacket(5)) + require.Equal(t, []interval{{Start: 4, End: 6}}, slices.Collect(hist.Backward())) + + // extend the range at the front + require.True(t, hist.ReceivedPacket(3)) + require.Equal(t, []interval{{Start: 3, End: 6}}, slices.Collect(hist.Backward())) +} + +func TestReceivedPacketHistoryRanges(t *testing.T) { + hist := newReceivedPacketHistory() + require.Equal(t, protocol.InvalidPacketNumber, hist.HighestMissingUpTo(1000)) + + require.True(t, hist.ReceivedPacket(4)) + require.Equal(t, protocol.PacketNumber(3), hist.HighestMissingUpTo(1000)) + require.Equal(t, protocol.PacketNumber(3), hist.HighestMissingUpTo(4)) + require.Equal(t, protocol.PacketNumber(3), hist.HighestMissingUpTo(3)) + require.Equal(t, protocol.PacketNumber(2), hist.HighestMissingUpTo(2)) + require.True(t, hist.ReceivedPacket(10)) + require.Equal(t, protocol.PacketNumber(9), hist.HighestMissingUpTo(1000)) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 4, End: 4}, + }, slices.Collect(hist.Backward())) + + // create a new range in the middle + require.True(t, hist.ReceivedPacket(7)) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 7, End: 7}, + {Start: 4, End: 4}, + }, slices.Collect(hist.Backward())) + + // create a new range at the front + require.True(t, hist.ReceivedPacket(1)) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 7, End: 7}, + {Start: 4, End: 4}, + {Start: 1, End: 1}, + }, slices.Collect(hist.Backward())) + + // extend an existing range at the end + require.True(t, hist.ReceivedPacket(8)) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 7, End: 8}, + {Start: 4, End: 4}, + {Start: 1, End: 1}, + }, slices.Collect(hist.Backward())) + + // extend an existing range at the front + require.True(t, hist.ReceivedPacket(6)) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 6, End: 8}, + {Start: 4, End: 4}, + {Start: 1, End: 1}, + }, slices.Collect(hist.Backward())) + + // close a range + require.True(t, hist.ReceivedPacket(9)) + require.Equal(t, []interval{ + {Start: 6, End: 10}, + {Start: 4, End: 4}, + {Start: 1, End: 1}, + }, slices.Collect(hist.Backward())) +} + +func TestReceivedPacketHistoryMaxNumAckRanges(t *testing.T) { + hist := newReceivedPacketHistory() + + for i := range protocol.MaxNumAckRanges { + require.True(t, hist.ReceivedPacket(protocol.PacketNumber(2*i))) + } + require.Len(t, hist.ranges, protocol.MaxNumAckRanges) + require.Equal(t, interval{Start: 0, End: 0}, hist.ranges[0]) + + hist.ReceivedPacket(2*protocol.MaxNumAckRanges + 1000) + // check that the oldest ACK range was deleted + require.Len(t, hist.ranges, protocol.MaxNumAckRanges) + require.Equal(t, interval{Start: 2, End: 2}, hist.ranges[0]) +} + +func TestReceivedPacketHistoryDeleteBelow(t *testing.T) { + hist := newReceivedPacketHistory() + + hist.DeleteBelow(2) + require.Empty(t, slices.Collect(hist.Backward())) + + require.True(t, hist.ReceivedPacket(2)) + require.True(t, hist.ReceivedPacket(4)) + require.True(t, hist.ReceivedPacket(5)) + require.True(t, hist.ReceivedPacket(6)) + require.True(t, hist.ReceivedPacket(10)) + + require.Equal(t, protocol.PacketNumber(3), hist.HighestMissingUpTo(6)) + hist.DeleteBelow(6) + require.Equal(t, protocol.InvalidPacketNumber, hist.HighestMissingUpTo(6)) + require.Equal(t, protocol.PacketNumber(9), hist.HighestMissingUpTo(10)) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 6, End: 6}, + }, slices.Collect(hist.Backward())) + + // deleting from an existing range + require.True(t, hist.ReceivedPacket(7)) + require.True(t, hist.ReceivedPacket(8)) + hist.DeleteBelow(7) + require.Equal(t, []interval{ + {Start: 10, End: 10}, + {Start: 7, End: 8}, + }, slices.Collect(hist.Backward())) + + // keep a one-packet range + hist.DeleteBelow(10) + require.Equal(t, []interval{{Start: 10, End: 10}}, slices.Collect(hist.Backward())) + + // delayed packets below deleted ranges are ignored + require.False(t, hist.ReceivedPacket(5)) + require.Equal(t, []interval{{Start: 10, End: 10}}, slices.Collect(hist.Backward())) +} + +func TestReceivedPacketHistoryDuplicateDetection(t *testing.T) { + hist := newReceivedPacketHistory() + + require.False(t, hist.IsPotentiallyDuplicate(5)) + + require.True(t, hist.ReceivedPacket(4)) + require.True(t, hist.ReceivedPacket(5)) + require.True(t, hist.ReceivedPacket(6)) + require.True(t, hist.ReceivedPacket(8)) + require.True(t, hist.ReceivedPacket(9)) + + require.False(t, hist.IsPotentiallyDuplicate(3)) + require.True(t, hist.IsPotentiallyDuplicate(4)) + require.True(t, hist.IsPotentiallyDuplicate(5)) + require.True(t, hist.IsPotentiallyDuplicate(6)) + require.False(t, hist.IsPotentiallyDuplicate(7)) + require.True(t, hist.IsPotentiallyDuplicate(8)) + require.True(t, hist.IsPotentiallyDuplicate(9)) + require.False(t, hist.IsPotentiallyDuplicate(10)) + + // delete and check for potential duplicates + hist.DeleteBelow(8) + require.True(t, hist.IsPotentiallyDuplicate(7)) + require.True(t, hist.IsPotentiallyDuplicate(8)) + require.True(t, hist.IsPotentiallyDuplicate(9)) + require.False(t, hist.IsPotentiallyDuplicate(10)) +} + +func TestReceivedPacketHistoryRandomized(t *testing.T) { + hist := newReceivedPacketHistory() + packets := make(map[protocol.PacketNumber]struct{}) + const num = 2 * protocol.MaxNumAckRanges + numLostPackets := rand.IntN(protocol.MaxNumAckRanges) + numRcvdPackets := num - numLostPackets + + for i := range num { + packets[protocol.PacketNumber(i)] = struct{}{} + } + lostPackets := make([]protocol.PacketNumber, 0, numLostPackets) + for len(lostPackets) < numLostPackets { + p := protocol.PacketNumber(rand.IntN(num - 1)) // lose a random packet, but not the last one + if _, ok := packets[p]; ok { + lostPackets = append(lostPackets, p) + delete(packets, p) + } + } + slices.Sort(lostPackets) + t.Logf("Losing packets: %v", lostPackets) + + ordered := make([]protocol.PacketNumber, 0, numRcvdPackets) + for p := range packets { + ordered = append(ordered, p) + } + rand.Shuffle(len(ordered), func(i, j int) { ordered[i], ordered[j] = ordered[j], ordered[i] }) + + t.Logf("Receiving packets: %v", ordered) + for i, p := range ordered { + require.True(t, hist.ReceivedPacket(p)) + // sometimes receive a duplicate + if i > 0 && rand.Int()%5 == 0 { + require.False(t, hist.ReceivedPacket(ordered[rand.IntN(i)])) + } + } + var counter int + ackRanges := slices.Collect(hist.Backward()) + t.Logf("ACK ranges: %v", ackRanges) + require.LessOrEqual(t, len(ackRanges), numLostPackets+1) + for _, ackRange := range ackRanges { + for p := ackRange.Start; p <= ackRange.End; p++ { + counter++ + require.Contains(t, packets, p) + } + } + require.Equal(t, numRcvdPackets, counter) + + deletedBelow := protocol.PacketNumber(rand.IntN(num * 2 / 3)) + t.Logf("Deleting below %d", deletedBelow) + hist.DeleteBelow(deletedBelow) + for pn := range protocol.PacketNumber(num) { + if pn < deletedBelow { + require.Equal(t, protocol.InvalidPacketNumber, hist.HighestMissingUpTo(pn)) + continue + } + expected := protocol.InvalidPacketNumber + for _, lost := range lostPackets { + if lost < deletedBelow { + continue + } + if lost > pn { + break + } + expected = lost + } + hm := hist.HighestMissingUpTo(pn) + require.Equalf(t, expected, hm, "highest missing up to %d: %d", pn, hm) + } +} + +func BenchmarkHistoryReceiveSequentialPackets(b *testing.B) { + hist := newReceivedPacketHistory() + var pn protocol.PacketNumber + for b.Loop() { + hist.ReceivedPacket(pn) + pn++ + } +} + +// Packets are received sequentially, with occasional gaps +func BenchmarkHistoryReceiveCommonCase(b *testing.B) { + hist := newReceivedPacketHistory() + var pn protocol.PacketNumber + for b.Loop() { + hist.ReceivedPacket(pn) + pn++ + if pn%2000 == 0 { + pn += 4 + } + } +} + +func BenchmarkHistoryReceiveSequentialPacketsWithGaps(b *testing.B) { + hist := newReceivedPacketHistory() + var pn protocol.PacketNumber + for b.Loop() { + hist.ReceivedPacket(pn) + pn += 2 + } +} + +func BenchmarkHistoryReceiveReversePacketsWithGaps(b *testing.B) { + hist := newReceivedPacketHistory() + for i := 0; i < b.N; i++ { + hist.ReceivedPacket(protocol.PacketNumber(2 * (b.N - i))) + } +} + +func BenchmarkHistoryIsDuplicate(b *testing.B) { + b.ReportAllocs() + hist := newReceivedPacketHistory() + var pn protocol.PacketNumber + for range protocol.MaxNumAckRanges { + for range 5 { + hist.ReceivedPacket(pn) + pn++ + } + pn += 5 // create a gap + } + + var p protocol.PacketNumber + for b.Loop() { + hist.IsPotentiallyDuplicate(p % pn) + p++ + } +} diff --git a/third_party/quic-go/internal/ackhandler/received_packet_tracker.go b/third_party/quic-go/internal/ackhandler/received_packet_tracker.go new file mode 100644 index 0000000..54c047b --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/received_packet_tracker.go @@ -0,0 +1,228 @@ +package ackhandler + +import ( + "fmt" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" +) + +const reorderingThreshold = 1 + +// The receivedPacketTracker tracks packets for the Initial and Handshake packet number space. +// Every received packet is acknowledged immediately. +type receivedPacketTracker struct { + ect0, ect1, ecnce uint64 + + packetHistory receivedPacketHistory + + lastAck *wire.AckFrame + hasNewAck bool // true as soon as we received an ack-eliciting new packet +} + +func newReceivedPacketTracker() *receivedPacketTracker { + return &receivedPacketTracker{packetHistory: *newReceivedPacketHistory()} +} + +func (h *receivedPacketTracker) ReceivedPacket(pn protocol.PacketNumber, ecn protocol.ECN, ackEliciting bool) error { + if isNew := h.packetHistory.ReceivedPacket(pn); !isNew { + return fmt.Errorf("receivedPacketTracker BUG: ReceivedPacket called for old / duplicate packet %d", pn) + } + + //nolint:exhaustive // Only need to count ECT(0), ECT(1) and ECN-CE. + switch ecn { + case protocol.ECT0: + h.ect0++ + case protocol.ECT1: + h.ect1++ + case protocol.ECNCE: + h.ecnce++ + } + if !ackEliciting { + return nil + } + h.hasNewAck = true + return nil +} + +func (h *receivedPacketTracker) GetAckFrame() *wire.AckFrame { + if !h.hasNewAck { + return nil + } + + // This function always returns the same ACK frame struct, filled with the most recent values. + ack := h.lastAck + if ack == nil { + ack = &wire.AckFrame{} + } + ack.Reset() + ack.ECT0 = h.ect0 + ack.ECT1 = h.ect1 + ack.ECNCE = h.ecnce + for r := range h.packetHistory.Backward() { + ack.AckRanges = append(ack.AckRanges, wire.AckRange{Smallest: r.Start, Largest: r.End}) + } + + h.lastAck = ack + h.hasNewAck = false + return ack +} + +func (h *receivedPacketTracker) IsPotentiallyDuplicate(pn protocol.PacketNumber) bool { + return h.packetHistory.IsPotentiallyDuplicate(pn) +} + +// number of ack-eliciting packets received before sending an ACK +const packetsBeforeAck = 2 + +// The appDataReceivedPacketTracker tracks packets received in the Application Data packet number space. +// It waits until at least 2 packets were received before queueing an ACK, or until the max_ack_delay was reached. +type appDataReceivedPacketTracker struct { + receivedPacketTracker + + largestObservedRcvdTime monotime.Time + + largestObserved protocol.PacketNumber + ignoreBelow protocol.PacketNumber + + maxAckDelay time.Duration + ackQueued bool // true if we need send a new ACK + + ackElicitingPacketsReceivedSinceLastAck int + ackAlarm monotime.Time + + logger utils.Logger +} + +func newAppDataReceivedPacketTracker(logger utils.Logger) *appDataReceivedPacketTracker { + h := &appDataReceivedPacketTracker{ + receivedPacketTracker: *newReceivedPacketTracker(), + maxAckDelay: protocol.MaxAckDelay, + logger: logger, + } + return h +} + +func (h *appDataReceivedPacketTracker) ReceivedPacket(pn protocol.PacketNumber, ecn protocol.ECN, rcvTime monotime.Time, ackEliciting bool) error { + if err := h.receivedPacketTracker.ReceivedPacket(pn, ecn, ackEliciting); err != nil { + return err + } + if pn >= h.largestObserved { + h.largestObserved = pn + h.largestObservedRcvdTime = rcvTime + } + if !ackEliciting { + return nil + } + h.ackElicitingPacketsReceivedSinceLastAck++ + isMissing := h.isMissing(pn) + if !h.ackQueued && h.shouldQueueACK(pn, ecn, isMissing) { + h.ackQueued = true + h.ackAlarm = 0 // cancel the ack alarm + } + if !h.ackQueued { + // No ACK queued, but we'll need to acknowledge the packet after max_ack_delay. + h.ackAlarm = rcvTime.Add(h.maxAckDelay) + if h.logger.Debug() { + h.logger.Debugf("\tSetting ACK timer to max ack delay: %s", h.maxAckDelay) + } + } + return nil +} + +// IgnoreBelow sets a lower limit for acknowledging packets. +// Packets with packet numbers smaller than p will not be acked. +func (h *appDataReceivedPacketTracker) IgnoreBelow(pn protocol.PacketNumber) { + if pn <= h.ignoreBelow { + return + } + h.ignoreBelow = pn + h.packetHistory.DeleteBelow(pn) + if h.logger.Debug() { + h.logger.Debugf("\tIgnoring all packets below %d.", pn) + } +} + +// isMissing says if a packet was reported missing in the last ACK. +func (h *appDataReceivedPacketTracker) isMissing(p protocol.PacketNumber) bool { + if h.lastAck == nil || p < h.ignoreBelow { + return false + } + return p < h.lastAck.LargestAcked() && !h.lastAck.AcksPacket(p) +} + +func (h *appDataReceivedPacketTracker) hasNewMissingPackets() bool { + if h.lastAck == nil { + return false + } + if h.largestObserved < reorderingThreshold { + return false + } + highestMissing := h.packetHistory.HighestMissingUpTo(h.largestObserved - reorderingThreshold) + if highestMissing == protocol.InvalidPacketNumber { + return false + } + if highestMissing < h.lastAck.LargestAcked() { + // the packet was already reported missing in the last ACK + return false + } + return highestMissing > h.lastAck.LargestAcked()-reorderingThreshold +} + +func (h *appDataReceivedPacketTracker) shouldQueueACK(pn protocol.PacketNumber, ecn protocol.ECN, wasMissing bool) bool { + // Send an ACK if this packet was reported missing in an ACK sent before. + // Ack decimation with reordering relies on the timer to send an ACK, but if + // missing packets we reported in the previous ACK, send an ACK immediately. + if wasMissing { + if h.logger.Debug() { + h.logger.Debugf("\tQueueing ACK because packet %d was missing before.", pn) + } + return true + } + + // send an ACK every 2 ack-eliciting packets + if h.ackElicitingPacketsReceivedSinceLastAck >= packetsBeforeAck { + if h.logger.Debug() { + h.logger.Debugf("\tQueueing ACK because packet %d packets were received after the last ACK (using initial threshold: %d).", h.ackElicitingPacketsReceivedSinceLastAck, packetsBeforeAck) + } + return true + } + + // queue an ACK if there are new missing packets to report + if h.hasNewMissingPackets() { + h.logger.Debugf("\tQueuing ACK because there's a new missing packet to report.") + return true + } + + // queue an ACK if the packet was ECN-CE marked + if ecn == protocol.ECNCE { + h.logger.Debugf("\tQueuing ACK because the packet was ECN-CE marked.") + return true + } + return false +} + +func (h *appDataReceivedPacketTracker) GetAckFrame(now monotime.Time, onlyIfQueued bool) *wire.AckFrame { + if onlyIfQueued && !h.ackQueued { + if h.ackAlarm.IsZero() || h.ackAlarm.After(now) { + return nil + } + if h.logger.Debug() && !h.ackAlarm.IsZero() { + h.logger.Debugf("Sending ACK because the ACK timer expired.") + } + } + ack := h.receivedPacketTracker.GetAckFrame() + if ack == nil { + return nil + } + ack.DelayTime = max(0, now.Sub(h.largestObservedRcvdTime)) + h.ackQueued = false + h.ackAlarm = 0 + h.ackElicitingPacketsReceivedSinceLastAck = 0 + return ack +} + +func (h *appDataReceivedPacketTracker) GetAlarmTimeout() monotime.Time { return h.ackAlarm } diff --git a/third_party/quic-go/internal/ackhandler/received_packet_tracker_test.go b/third_party/quic-go/internal/ackhandler/received_packet_tracker_test.go new file mode 100644 index 0000000..b590333 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/received_packet_tracker_test.go @@ -0,0 +1,188 @@ +package ackhandler + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestReceivedPacketTrackerGenerateACKs(t *testing.T) { + tracker := newReceivedPacketTracker() + + require.NoError(t, tracker.ReceivedPacket(protocol.PacketNumber(3), protocol.ECNNon, true)) + ack := tracker.GetAckFrame() + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 3, Largest: 3}}, ack.AckRanges) + require.Zero(t, ack.DelayTime) + + require.NoError(t, tracker.ReceivedPacket(protocol.PacketNumber(4), protocol.ECNNon, true)) + ack = tracker.GetAckFrame() + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 3, Largest: 4}}, ack.AckRanges) + require.Zero(t, ack.DelayTime) + + require.NoError(t, tracker.ReceivedPacket(protocol.PacketNumber(1), protocol.ECNNon, true)) + ack = tracker.GetAckFrame() + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{ + {Smallest: 3, Largest: 4}, + {Smallest: 1, Largest: 1}, + }, ack.AckRanges) + require.Zero(t, ack.DelayTime) + + // non-ack-eliciting packets don't trigger ACKs + require.NoError(t, tracker.ReceivedPacket(protocol.PacketNumber(10), protocol.ECNNon, false)) + require.Nil(t, tracker.GetAckFrame()) + + require.NoError(t, tracker.ReceivedPacket(protocol.PacketNumber(11), protocol.ECNNon, true)) + ack = tracker.GetAckFrame() + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{ + {Smallest: 10, Largest: 11}, + {Smallest: 3, Largest: 4}, + {Smallest: 1, Largest: 1}, + }, ack.AckRanges) +} + +func TestAppDataReceivedPacketTrackerECN(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + + require.NoError(t, tr.ReceivedPacket(0, protocol.ECT0, monotime.Now(), true)) + pn := protocol.PacketNumber(1) + for range 2 { + require.NoError(t, tr.ReceivedPacket(pn, protocol.ECT1, monotime.Now(), true)) + pn++ + } + for range 3 { + require.NoError(t, tr.ReceivedPacket(pn, protocol.ECNCE, monotime.Now(), true)) + pn++ + } + ack := tr.GetAckFrame(monotime.Now(), false) + require.Equal(t, uint64(1), ack.ECT0) + require.Equal(t, uint64(2), ack.ECT1) + require.Equal(t, uint64(3), ack.ECNCE) +} + +func TestAppDataReceivedPacketTrackerAckEverySecondPacket(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + require.Nil(t, tr.GetAckFrame(monotime.Now(), true)) + + for p := protocol.PacketNumber(1); p <= 20; p++ { + require.NoError(t, tr.ReceivedPacket(p, protocol.ECNNon, monotime.Now(), true)) + switch p % 2 { + case 0: + require.NotNil(t, tr.GetAckFrame(monotime.Now(), true)) + case 1: + require.Nil(t, tr.GetAckFrame(monotime.Now(), true)) + } + } +} + +func TestAppDataReceivedPacketTrackerAlarmTimeout(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + + now := monotime.Now() + require.NoError(t, tr.ReceivedPacket(1, protocol.ECNNon, now, false)) + require.Nil(t, tr.GetAckFrame(monotime.Now(), true)) + require.Zero(t, tr.GetAlarmTimeout()) + + rcvTime := now.Add(10 * time.Millisecond) + require.NoError(t, tr.ReceivedPacket(2, protocol.ECNNon, rcvTime, true)) + require.Equal(t, rcvTime.Add(protocol.MaxAckDelay), tr.GetAlarmTimeout()) + require.Nil(t, tr.GetAckFrame(monotime.Now(), true)) + + // no timeout after the ACK has been dequeued + require.NotNil(t, tr.GetAckFrame(monotime.Now(), false)) + require.Zero(t, tr.GetAlarmTimeout()) +} + +func TestAppDataReceivedPacketTrackerQueuesECNCE(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + + require.NoError(t, tr.ReceivedPacket(1, protocol.ECNCE, monotime.Now(), true)) + ack := tr.GetAckFrame(monotime.Now(), true) + require.NotNil(t, ack) + require.Equal(t, protocol.PacketNumber(1), ack.LargestAcked()) + require.EqualValues(t, 1, ack.ECNCE) +} + +func TestAppDataReceivedPacketTrackerMissingPackets(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + + now := monotime.Now() + require.NoError(t, tr.ReceivedPacket(0, protocol.ECNNon, now, true)) + require.Nil(t, tr.GetAckFrame(now, true)) + + require.NoError(t, tr.ReceivedPacket(5, protocol.ECNNon, now, true)) + ack := tr.GetAckFrame(now, true) // ACK: 0 and 5, missing: 1, 2, 3, 4 + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 5, Largest: 5}, {Smallest: 0, Largest: 0}}, ack.AckRanges) + + // now receive one of the missing packets + require.NoError(t, tr.ReceivedPacket(3, protocol.ECNNon, now, true)) + ack = tr.GetAckFrame(now, true) + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{ + {Smallest: 5, Largest: 5}, + {Smallest: 3, Largest: 3}, + {Smallest: 0, Largest: 0}, + }, ack.AckRanges) + + require.NoError(t, tr.ReceivedPacket(6, protocol.ECNNon, now, true)) + require.Nil(t, tr.GetAckFrame(now, true)) + require.NoError(t, tr.ReceivedPacket(8, protocol.ECNNon, now, true)) + require.NotNil(t, tr.GetAckFrame(now, true)) +} + +func TestAppDataReceivedPacketTrackerDelayTime(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + + now := monotime.Now() + require.NoError(t, tr.ReceivedPacket(1, protocol.ECNNon, now, true)) + require.NoError(t, tr.ReceivedPacket(2, protocol.ECNNon, now.Add(-1337*time.Millisecond), true)) + ack := tr.GetAckFrame(now, true) + require.NotNil(t, ack) + require.Equal(t, 1337*time.Millisecond, ack.DelayTime) + + // don't use a negative delay time + require.NoError(t, tr.ReceivedPacket(3, protocol.ECNNon, now.Add(time.Hour), true)) + ack = tr.GetAckFrame(now, false) + require.NotNil(t, ack) + require.Zero(t, ack.DelayTime) +} + +func TestAppDataReceivedPacketTrackerIgnoreBelow(t *testing.T) { + tr := newAppDataReceivedPacketTracker(utils.DefaultLogger) + + tr.IgnoreBelow(4) + // check that packets below 7 are considered duplicates + require.True(t, tr.IsPotentiallyDuplicate(3)) + require.False(t, tr.IsPotentiallyDuplicate(4)) + + for i := 5; i <= 10; i++ { + require.NoError(t, tr.ReceivedPacket(protocol.PacketNumber(i), protocol.ECNNon, monotime.Now(), true)) + } + ack := tr.GetAckFrame(monotime.Now(), true) + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 5, Largest: 10}}, ack.AckRanges) + + tr.IgnoreBelow(7) + + require.NoError(t, tr.ReceivedPacket(11, protocol.ECNNon, monotime.Now(), true)) + require.NoError(t, tr.ReceivedPacket(12, protocol.ECNNon, monotime.Now(), true)) + ack = tr.GetAckFrame(monotime.Now(), true) + require.NotNil(t, ack) + require.Equal(t, []wire.AckRange{{Smallest: 7, Largest: 12}}, ack.AckRanges) + + // make sure that old packets are not accepted + require.ErrorContains(t, + tr.ReceivedPacket(4, protocol.ECNNon, monotime.Now(), true), + "receivedPacketTracker BUG: ReceivedPacket called for old / duplicate packet 4", + ) +} diff --git a/third_party/quic-go/internal/ackhandler/send_mode.go b/third_party/quic-go/internal/ackhandler/send_mode.go new file mode 100644 index 0000000..c03f3a6 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/send_mode.go @@ -0,0 +1,46 @@ +package ackhandler + +import "fmt" + +// The SendMode says what kind of packets can be sent. +type SendMode uint8 + +const ( + // SendNone means that no packets should be sent + SendNone SendMode = iota + // SendAck means an ACK-only packet should be sent + SendAck + // SendPTOInitial means that an Initial probe packet should be sent + SendPTOInitial + // SendPTOHandshake means that a Handshake probe packet should be sent + SendPTOHandshake + // SendPTOAppData means that an Application data probe packet should be sent + SendPTOAppData + // SendPacingLimited means that the pacer doesn't allow sending of a packet right now, + // but will do in a little while. + // The timestamp when sending is allowed again can be obtained via the SentPacketHandler.TimeUntilSend. + SendPacingLimited + // SendAny means that any packet should be sent + SendAny +) + +func (s SendMode) String() string { + switch s { + case SendNone: + return "none" + case SendAck: + return "ack" + case SendPTOInitial: + return "pto (Initial)" + case SendPTOHandshake: + return "pto (Handshake)" + case SendPTOAppData: + return "pto (Application Data)" + case SendAny: + return "any" + case SendPacingLimited: + return "pacing limited" + default: + return fmt.Sprintf("invalid send mode: %d", s) + } +} diff --git a/third_party/quic-go/internal/ackhandler/send_mode_test.go b/third_party/quic-go/internal/ackhandler/send_mode_test.go new file mode 100644 index 0000000..0159ba0 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/send_mode_test.go @@ -0,0 +1,18 @@ +package ackhandler + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestSendModeStringer(t *testing.T) { + require.Equal(t, "none", SendNone.String()) + require.Equal(t, "any", SendAny.String()) + require.Equal(t, "pacing limited", SendPacingLimited.String()) + require.Equal(t, "ack", SendAck.String()) + require.Equal(t, "pto (Initial)", SendPTOInitial.String()) + require.Equal(t, "pto (Handshake)", SendPTOHandshake.String()) + require.Equal(t, "pto (Application Data)", SendPTOAppData.String()) + require.Equal(t, "invalid send mode: 123", SendMode(123).String()) +} diff --git a/third_party/quic-go/internal/ackhandler/sent_packet_handler.go b/third_party/quic-go/internal/ackhandler/sent_packet_handler.go new file mode 100644 index 0000000..16d9acc --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/sent_packet_handler.go @@ -0,0 +1,1240 @@ +package ackhandler + +import ( + "errors" + "fmt" + "sync" + "time" + + congestionExt "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/internal/congestion" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +const ( + // Maximum reordering in time space before time based loss detection considers a packet lost. + // Specified as an RTT multiplier. + timeThreshold = 9.0 / 8 + // Maximum reordering in packets before packet threshold loss detection considers a packet lost. + packetThreshold = 3 + // Before validating the client's address, the server won't send more than 3x bytes than it received. + amplificationFactor = 3 + // We use Retry packets to derive an RTT estimate. Make sure we don't set the RTT to a super low value yet. + minRTTAfterRetry = 5 * time.Millisecond + // The PTO duration uses exponential backoff, but is truncated to a maximum value, as allowed by RFC 8961, section 4.4. + maxPTODuration = 60 * time.Second +) + +// Path probe packets are declared lost after this time. +const pathProbePacketLossTimeout = time.Second + +// minDatagramRoomForSizing is the least room left in a datagram that is still +// used to size the next packet number. The imitated client needs the room to +// fit a header before it caps a packet to it; below that it keeps the full +// datagram size. A header size is not a constant, so this approximates it. +const minDatagramRoomForSizing protocol.ByteCount = 40 + +type packetNumberSpace struct { + history sentPacketHistory + pns packetNumberGenerator + + lossTime monotime.Time + lastAckElicitingPacketTime monotime.Time + + largestAcked protocol.PacketNumber + largestSent protocol.PacketNumber +} + +func newPacketNumberSpace(initialPN protocol.PacketNumber, isAppData bool) *packetNumberSpace { + var pns packetNumberGenerator + if isAppData { + pns = newSkippingPacketNumberGenerator(initialPN, protocol.SkipPacketInitialPeriod, protocol.SkipPacketMaxPeriod) + } else { + pns = newSequentialPacketNumberGenerator(initialPN) + } + return &packetNumberSpace{ + history: *newSentPacketHistory(isAppData), + pns: pns, + largestSent: protocol.InvalidPacketNumber, + largestAcked: protocol.InvalidPacketNumber, + } +} + +type alarmTimer struct { + Time monotime.Time + TimerType qlog.TimerType + EncryptionLevel protocol.EncryptionLevel +} + +type sentPacketHandler struct { + initialPackets *packetNumberSpace + handshakePackets *packetNumberSpace + appDataPackets *packetNumberSpace + lostPackets lostPacketTracker // only for application-data packet number space + // send time of the largest acknowledged packet, across all packet number spaces + largestAckedTime monotime.Time + + // Do we know that the peer completed address validation yet? + // Always true for the server. + peerCompletedAddressValidation bool + bytesReceived protocol.ByteCount + bytesSent protocol.ByteCount + // Have we validated the peer's address yet? + // Always true for the client. + peerAddressValidated bool + + handshakeConfirmed bool + + ignorePacketsBelow func(protocol.PacketNumber) + + ackedPackets []packetWithPacketNumber // to avoid allocations in detectAndRemoveAckedPackets + ackedPacketsInfo []congestionExt.AckedPacketInfo + lostPacketsInfo []congestionExt.LostPacketInfo + + bytesInFlight protocol.ByteCount + + congestion congestion.SendAlgorithmWithDebugInfos + congestionMutex sync.RWMutex + rttStats *utils.RTTStats + connStats *utils.ConnectionStats + + // The number of times a PTO has been sent without receiving an ack. + ptoCount uint32 + ptoMode SendMode + // The number of PTO probe packets that should be sent. + // Only applies to the application-data packet number space. + numProbesToSend int + + // The alarm timeout + alarm alarmTimer + + enableECN bool + ecnTracker ecnHandler + + perspective protocol.Perspective + + // shortPacketNumbers allows single-byte packet numbers, as Chrome uses. + shortPacketNumbers bool + // maxDatagramSize converts the congestion window into a packet count, which + // is what the packet number length depends on when shortPacketNumbers is set. + maxDatagramSize protocol.ByteCount + // lastDatagramPadding is the room left in the datagram that was packed last, + // which is what the imitated client measures the window against instead. + lastDatagramPadding protocol.ByteCount + + qlogger qlogwriter.Recorder + lastMetrics qlog.MetricsUpdated + logger utils.Logger +} + +var _ SentPacketHandler = &sentPacketHandler{} + +// clientAddressValidated indicates whether the address was validated beforehand by an address validation token. +// If the address was validated, the amplification limit doesn't apply. It has no effect for a client. +func NewSentPacketHandler( + initialPN protocol.PacketNumber, + initialMaxDatagramSize protocol.ByteCount, + rttStats *utils.RTTStats, + connStats *utils.ConnectionStats, + clientAddressValidated bool, + enableECN bool, + ignorePacketsBelow func(protocol.PacketNumber), + pers protocol.Perspective, + shortPacketNumbers bool, + qlogger qlogwriter.Recorder, + logger utils.Logger, +) SentPacketHandler { + congestion := congestion.NewCubicSender( + congestion.DefaultClock{}, + rttStats, + connStats, + initialMaxDatagramSize, + true, // use Reno + qlogger, + ) + + h := &sentPacketHandler{ + shortPacketNumbers: shortPacketNumbers, + maxDatagramSize: initialMaxDatagramSize, + peerCompletedAddressValidation: pers == protocol.PerspectiveServer, + peerAddressValidated: pers == protocol.PerspectiveClient || clientAddressValidated, + initialPackets: newPacketNumberSpace(initialPN, false), + handshakePackets: newPacketNumberSpace(0, false), + appDataPackets: newPacketNumberSpace(0, true), + lostPackets: *newLostPacketTracker(64), + rttStats: rttStats, + connStats: connStats, + congestion: congestion, + ignorePacketsBelow: ignorePacketsBelow, + perspective: pers, + qlogger: qlogger, + logger: logger, + } + if enableECN { + h.enableECN = true + h.ecnTracker = newECNTracker(logger, qlogger) + } + return h +} + +func (h *sentPacketHandler) removeFromBytesInFlight(p *packet) { + if p.includedInBytesInFlight { + if p.Length > h.bytesInFlight { + panic("negative bytes_in_flight") + } + h.bytesInFlight -= p.Length + p.includedInBytesInFlight = false + } +} + +func (h *sentPacketHandler) DropPackets(encLevel protocol.EncryptionLevel, now monotime.Time) { + // The server won't await address validation after the handshake is confirmed. + // This applies even if we didn't receive an ACK for a Handshake packet. + if h.perspective == protocol.PerspectiveClient && encLevel == protocol.EncryptionHandshake { + h.peerCompletedAddressValidation = true + } + // remove outstanding packets from bytes_in_flight + if encLevel == protocol.EncryptionInitial || encLevel == protocol.EncryptionHandshake { + pnSpace := h.getPacketNumberSpace(encLevel) + // We might already have dropped this packet number space. + if pnSpace == nil { + return + } + for _, p := range pnSpace.history.Packets() { + h.removeFromBytesInFlight(p) + } + } + // drop the packet history + //nolint:exhaustive // Not every packet number space can be dropped. + switch encLevel { + case protocol.EncryptionInitial: + h.initialPackets = nil + case protocol.EncryptionHandshake: + // Dropping the handshake packet number space means that the handshake is confirmed, + // see section 4.9.2 of RFC 9001. + h.handshakeConfirmed = true + h.handshakePackets = nil + case protocol.Encryption0RTT: + // This function is only called when 0-RTT is rejected, + // and not when the client drops 0-RTT keys when the handshake completes. + // When 0-RTT is rejected, all application data sent so far becomes invalid. + // Delete the packets from the history and remove them from bytes_in_flight. + for pn, p := range h.appDataPackets.history.Packets() { + if p.EncryptionLevel != protocol.Encryption0RTT { + break + } + h.removeFromBytesInFlight(p) + h.appDataPackets.history.Remove(pn) + } + default: + panic(fmt.Sprintf("Cannot drop keys for encryption level %s", encLevel)) + } + if h.qlogger != nil && h.ptoCount != 0 { + h.qlogger.RecordEvent(qlog.PTOCountUpdated{PTOCount: 0}) + } + h.ptoCount = 0 + h.numProbesToSend = 0 + h.ptoMode = SendNone + h.setLossDetectionTimer(now) +} + +func (h *sentPacketHandler) ReceivedBytes(n protocol.ByteCount, t monotime.Time) { + h.connStats.BytesReceived.Add(uint64(n)) + wasAmplificationLimit := h.isAmplificationLimited() + h.bytesReceived += n + if wasAmplificationLimit && !h.isAmplificationLimited() { + h.setLossDetectionTimer(t) + } +} + +func (h *sentPacketHandler) ReceivedPacket(l protocol.EncryptionLevel, t monotime.Time) { + h.connStats.PacketsReceived.Add(1) + if h.perspective == protocol.PerspectiveServer && l == protocol.EncryptionHandshake && !h.peerAddressValidated { + h.peerAddressValidated = true + h.setLossDetectionTimer(t) + } +} + +func (h *sentPacketHandler) packetsInFlight() int { + packetsInFlight := h.appDataPackets.history.NumOutstanding() + if h.handshakePackets != nil { + packetsInFlight += h.handshakePackets.history.NumOutstanding() + } + if h.initialPackets != nil { + packetsInFlight += h.initialPackets.history.NumOutstanding() + } + return packetsInFlight +} + +func (h *sentPacketHandler) SentPacket( + t monotime.Time, + pn, largestAcked protocol.PacketNumber, + streamFrames []StreamFrame, + frames []Frame, + encLevel protocol.EncryptionLevel, + ecn protocol.ECN, + size protocol.ByteCount, + isPathMTUProbePacket bool, + isPathProbePacket bool, +) { + h.bytesSent += size + h.connStats.BytesSent.Add(uint64(size)) + h.connStats.PacketsSent.Add(1) + + pnSpace := h.getPacketNumberSpace(encLevel) + if h.logger.Debug() && (pnSpace.history.HasOutstandingPackets() || pnSpace.history.HasOutstandingPathProbes()) { + for p := max(0, pnSpace.largestSent+1); p < pn; p++ { + h.logger.Debugf("Skipping packet number %d", p) + } + } + + pnSpace.largestSent = pn + + p := getPacket() + p.SendTime = t + p.EncryptionLevel = encLevel + p.Length = size + p.Frames = frames + p.LargestAcked = largestAcked + p.StreamFrames = streamFrames + p.IsPathMTUProbePacket = isPathMTUProbePacket + p.isPathProbePacket = isPathProbePacket + isAckEliciting := p.IsAckEliciting() + + if isPathProbePacket { + pnSpace.history.SentPathProbePacket(pn, p) + h.setLossDetectionTimer(t) + return + } + if isAckEliciting { + pnSpace.lastAckElicitingPacketTime = t + h.bytesInFlight += size + p.includedInBytesInFlight = true + if h.numProbesToSend > 0 { + h.numProbesToSend-- + } + } + + cc := h.getCongestionControl() + cc.OnPacketSent(t, h.bytesInFlight, pn, size, isAckEliciting) + + if encLevel == protocol.Encryption1RTT && h.ecnTracker != nil { + h.ecnTracker.SentPacket(pn, ecn) + } + + pnSpace.history.SentPacket(pn, p) + if !isAckEliciting { + if !h.peerCompletedAddressValidation { + h.setLossDetectionTimer(t) + } + return + } + if h.qlogger != nil { + h.qlogMetricsUpdated() + } + h.setLossDetectionTimer(t) +} + +func (h *sentPacketHandler) qlogMetricsUpdated() { + var metricsUpdatedEvent qlog.MetricsUpdated + var updated bool + if h.rttStats.HasMeasurement() { + if h.lastMetrics.MinRTT != h.rttStats.MinRTT() { + metricsUpdatedEvent.MinRTT = h.rttStats.MinRTT() + h.lastMetrics.MinRTT = metricsUpdatedEvent.MinRTT + updated = true + } + if h.lastMetrics.SmoothedRTT != h.rttStats.SmoothedRTT() { + metricsUpdatedEvent.SmoothedRTT = h.rttStats.SmoothedRTT() + h.lastMetrics.SmoothedRTT = metricsUpdatedEvent.SmoothedRTT + updated = true + } + if h.lastMetrics.LatestRTT != h.rttStats.LatestRTT() { + metricsUpdatedEvent.LatestRTT = h.rttStats.LatestRTT() + h.lastMetrics.LatestRTT = metricsUpdatedEvent.LatestRTT + updated = true + } + if h.lastMetrics.RTTVariance != h.rttStats.MeanDeviation() { + metricsUpdatedEvent.RTTVariance = h.rttStats.MeanDeviation() + h.lastMetrics.RTTVariance = metricsUpdatedEvent.RTTVariance + updated = true + } + } + if cwnd := h.getCongestionControl().GetCongestionWindow(); h.lastMetrics.CongestionWindow != int(cwnd) { + metricsUpdatedEvent.CongestionWindow = int(cwnd) + h.lastMetrics.CongestionWindow = metricsUpdatedEvent.CongestionWindow + updated = true + } + if h.lastMetrics.BytesInFlight != int(h.bytesInFlight) { + metricsUpdatedEvent.BytesInFlight = int(h.bytesInFlight) + h.lastMetrics.BytesInFlight = metricsUpdatedEvent.BytesInFlight + updated = true + } + packetsInFlight := h.packetsInFlight() + if h.lastMetrics.PacketsInFlight != packetsInFlight { + metricsUpdatedEvent.PacketsInFlight = packetsInFlight + h.lastMetrics.PacketsInFlight = metricsUpdatedEvent.PacketsInFlight + updated = true + } + if updated { + h.qlogger.RecordEvent(metricsUpdatedEvent) + } +} + +func (h *sentPacketHandler) getPacketNumberSpace(encLevel protocol.EncryptionLevel) *packetNumberSpace { + switch encLevel { + case protocol.EncryptionInitial: + return h.initialPackets + case protocol.EncryptionHandshake: + return h.handshakePackets + case protocol.Encryption0RTT, protocol.Encryption1RTT: + return h.appDataPackets + default: + panic("invalid packet number space") + } +} + +func (h *sentPacketHandler) ReceivedAck(ack *wire.AckFrame, encLevel protocol.EncryptionLevel, rcvTime monotime.Time) (bool /* contained 1-RTT packet */, error) { + pnSpace := h.getPacketNumberSpace(encLevel) + + largestAcked := ack.LargestAcked() + if largestAcked > pnSpace.largestSent { + return false, &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "received ACK for an unsent packet", + } + } + + // Servers complete address validation when a protected packet is received. + if h.perspective == protocol.PerspectiveClient && !h.peerCompletedAddressValidation && + (encLevel == protocol.EncryptionHandshake || encLevel == protocol.Encryption1RTT) { + h.peerCompletedAddressValidation = true + h.logger.Debugf("Peer doesn't await address validation any longer.") + // Make sure that the timer is reset, even if this ACK doesn't acknowledge any (ack-eliciting) packets. + h.setLossDetectionTimer(rcvTime) + } + + priorInFlight := h.bytesInFlight + ackedPackets, hasAckEliciting, err := h.detectAndRemoveAckedPackets(ack, encLevel) + if err != nil || len(ackedPackets) == 0 { + return false, err + } + + cc := h.getCongestionControl() + + // update the RTT, if: + // * the largest acked is newly acknowledged, AND + // * at least one new ack-eliciting packet was acknowledged + if len(ackedPackets) > 0 { + if p := ackedPackets[len(ackedPackets)-1]; p.PacketNumber == ack.LargestAcked() && !p.isPathProbePacket && hasAckEliciting { + // don't use the ack delay for Initial and Handshake packets + var ackDelay time.Duration + if encLevel == protocol.Encryption1RTT { + ackDelay = min(ack.DelayTime, h.rttStats.MaxAckDelay()) + } + if h.largestAckedTime.IsZero() || !p.SendTime.Before(h.largestAckedTime) { + h.rttStats.UpdateRTT(rcvTime.Sub(p.SendTime), ackDelay) + if h.logger.Debug() { + h.logger.Debugf("\tupdated RTT: %s (σ: %s)", h.rttStats.SmoothedRTT(), h.rttStats.MeanDeviation()) + } + h.largestAckedTime = p.SendTime + } + cc.MaybeExitSlowStart() + } + } + + // Only inform the ECN tracker about new 1-RTT ACKs if the ACK increases the largest acked. + if encLevel == protocol.Encryption1RTT && h.ecnTracker != nil && largestAcked > pnSpace.largestAcked { + congested := h.ecnTracker.HandleNewlyAcked(ackedPackets, int64(ack.ECT0), int64(ack.ECT1), int64(ack.ECNCE)) + if congested { + cc.OnCongestionEvent(largestAcked, 0, priorInFlight) + } + } + + pnSpace.largestAcked = max(pnSpace.largestAcked, largestAcked) + + h.detectLostPackets(rcvTime, encLevel) + h.ackedPacketsInfo = h.ackedPacketsInfo[:0] + if encLevel == protocol.Encryption1RTT { + h.detectLostPathProbes(rcvTime) + } + var acked1RTTPacket bool + for _, p := range ackedPackets { + if p.includedInBytesInFlight { + cc.OnPacketAcked(p.PacketNumber, p.Length, priorInFlight, rcvTime) + h.ackedPacketsInfo = append(h.ackedPacketsInfo, congestionExt.AckedPacketInfo{ + PacketNumber: congestionExt.PacketNumber(p.PacketNumber), + BytesAcked: congestionExt.ByteCount(p.Length), + }) + } + if p.EncryptionLevel == protocol.Encryption1RTT { + acked1RTTPacket = true + } + h.removeFromBytesInFlight(p.packet) + if !p.isPathProbePacket { + putPacket(p.packet) + } + } + + if cex, ok := h.getCongestionControl().(congestion.SendAlgorithmEx); ok && + (len(h.ackedPacketsInfo) != 0 || len(h.lostPacketsInfo) != 0) { + cex.OnCongestionEventEx(priorInFlight, rcvTime, h.ackedPacketsInfo, h.lostPacketsInfo) + } + + // detect spurious losses for application data packets, if the ACK was not reordered + if encLevel == protocol.Encryption1RTT && largestAcked == pnSpace.largestAcked { + h.detectSpuriousLosses( + ack, + rcvTime.Add(-min(ack.DelayTime, h.rttStats.MaxAckDelay())), + ) + // clean up lost packet history + h.lostPackets.DeleteBefore(rcvTime.Add(-3 * h.rttStats.PTO(false))) + } + + // After this point, we must not use ackedPackets any longer! + // We've already returned the buffers. + ackedPackets = nil //nolint:ineffassign // This is just to be on the safe side. + clear(h.ackedPackets) // make sure the memory is released + h.ackedPackets = h.ackedPackets[:0] + h.ackedPacketsInfo = nil //nolint:ineffassign // This is just to be on the safe side. + h.lostPacketsInfo = nil //nolint:ineffassign // This is just to be on the safe side. + + // Reset the pto_count unless the client is unsure if the server has validated the client's address. + if h.peerCompletedAddressValidation { + if h.qlogger != nil && h.ptoCount != 0 { + h.qlogger.RecordEvent(qlog.PTOCountUpdated{PTOCount: 0}) + } + h.ptoCount = 0 + } + h.numProbesToSend = 0 + + if h.qlogger != nil { + h.qlogMetricsUpdated() + } + + h.setLossDetectionTimer(rcvTime) + return acked1RTTPacket, nil +} + +func (h *sentPacketHandler) detectSpuriousLosses(ack *wire.AckFrame, ackTime monotime.Time) { + var maxPacketReordering protocol.PacketNumber + var maxTimeReordering time.Duration + ackRangeIdx := len(ack.AckRanges) - 1 + var spuriousLosses []protocol.PacketNumber + for pn, sendTime := range h.lostPackets.All() { + ackRange := ack.AckRanges[ackRangeIdx] + for pn > ackRange.Largest { + // this should never happen, since detectSpuriousLosses is only called for ACKs that increase the largest acked + if ackRangeIdx == 0 { + break + } + ackRangeIdx-- + ackRange = ack.AckRanges[ackRangeIdx] + } + if pn < ackRange.Smallest { + continue + } + if pn <= ackRange.Largest { + packetReordering := h.appDataPackets.history.Difference(ack.LargestAcked(), pn) + timeReordering := ackTime.Sub(sendTime) + maxPacketReordering = max(maxPacketReordering, packetReordering) + maxTimeReordering = max(maxTimeReordering, timeReordering) + + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pn, + PacketReordering: uint64(packetReordering), + TimeReordering: timeReordering, + }) + } + spuriousLosses = append(spuriousLosses, pn) + } + } + for _, pn := range spuriousLosses { + h.lostPackets.Delete(pn) + } +} + +// Packets are returned in ascending packet number order. +func (h *sentPacketHandler) detectAndRemoveAckedPackets( + ack *wire.AckFrame, + encLevel protocol.EncryptionLevel, +) (_ []packetWithPacketNumber, hasAckEliciting bool, _ error) { + if len(h.ackedPackets) > 0 { + return nil, false, errors.New("ackhandler BUG: ackedPackets slice not empty") + } + + pnSpace := h.getPacketNumberSpace(encLevel) + + if encLevel == protocol.Encryption1RTT { + for p := range pnSpace.history.SkippedPackets() { + if ack.AcksPacket(p) { + return nil, false, &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: fmt.Sprintf("received an ACK for skipped packet number: %d (%s)", p, encLevel), + } + } + } + } + + var ackRangeIndex int + lowestAcked := ack.LowestAcked() + largestAcked := ack.LargestAcked() + for pn, p := range pnSpace.history.Packets() { + // ignore packets below the lowest acked + if pn < lowestAcked { + continue + } + if pn > largestAcked { + break + } + + if ack.HasMissingRanges() { + ackRange := ack.AckRanges[len(ack.AckRanges)-1-ackRangeIndex] + + for pn > ackRange.Largest && ackRangeIndex < len(ack.AckRanges)-1 { + ackRangeIndex++ + ackRange = ack.AckRanges[len(ack.AckRanges)-1-ackRangeIndex] + } + + if pn < ackRange.Smallest { // packet not contained in ACK range + continue + } + if pn > ackRange.Largest { + return nil, false, fmt.Errorf("BUG: ackhandler would have acked wrong packet %d, while evaluating range %d -> %d", pn, ackRange.Smallest, ackRange.Largest) + } + } + if p.isPathProbePacket { + probePacket := pnSpace.history.RemovePathProbe(pn) + // the probe packet might already have been declared lost + if probePacket != nil { + h.ackedPackets = append(h.ackedPackets, packetWithPacketNumber{PacketNumber: pn, packet: probePacket}) + } + continue + } + if p.IsAckEliciting() { + hasAckEliciting = true + } + h.ackedPackets = append(h.ackedPackets, packetWithPacketNumber{PacketNumber: pn, packet: p}) + } + if h.logger.Debug() && len(h.ackedPackets) > 0 { + pns := make([]protocol.PacketNumber, len(h.ackedPackets)) + for i, p := range h.ackedPackets { + pns[i] = p.PacketNumber + } + h.logger.Debugf("\tnewly acked packets (%d): %d", len(pns), pns) + } + + for _, p := range h.ackedPackets { + if p.LargestAcked != protocol.InvalidPacketNumber && encLevel == protocol.Encryption1RTT && h.ignorePacketsBelow != nil { + h.ignorePacketsBelow(p.LargestAcked + 1) + } + + for _, f := range p.Frames { + if f.Handler != nil { + f.Handler.OnAcked(f.Frame) + } + } + for _, f := range p.StreamFrames { + if f.Handler != nil { + f.Handler.OnAcked(f.Frame) + } + } + if err := pnSpace.history.Remove(p.PacketNumber); err != nil { + return nil, false, err + } + } + // TODO: add support for the transport:packets_acked qlog event + return h.ackedPackets, hasAckEliciting, nil +} + +func (h *sentPacketHandler) getLossTimeAndSpace() (monotime.Time, protocol.EncryptionLevel) { + var encLevel protocol.EncryptionLevel + var lossTime monotime.Time + + if h.initialPackets != nil { + lossTime = h.initialPackets.lossTime + encLevel = protocol.EncryptionInitial + } + if h.handshakePackets != nil && (lossTime.IsZero() || (!h.handshakePackets.lossTime.IsZero() && h.handshakePackets.lossTime.Before(lossTime))) { + lossTime = h.handshakePackets.lossTime + encLevel = protocol.EncryptionHandshake + } + if lossTime.IsZero() || (!h.appDataPackets.lossTime.IsZero() && h.appDataPackets.lossTime.Before(lossTime)) { + lossTime = h.appDataPackets.lossTime + encLevel = protocol.Encryption1RTT + } + return lossTime, encLevel +} + +func (h *sentPacketHandler) getScaledPTO(includeMaxAckDelay bool) time.Duration { + pto := h.rttStats.PTO(includeMaxAckDelay) << h.ptoCount + if pto > maxPTODuration || pto <= 0 { + return maxPTODuration + } + return pto +} + +// same logic as getLossTimeAndSpace, but for lastAckElicitingPacketTime instead of lossTime +func (h *sentPacketHandler) getPTOTimeAndSpace(now monotime.Time) (pto monotime.Time, encLevel protocol.EncryptionLevel) { + // We only send application data probe packets once the handshake is confirmed, + // because before that, we don't have the keys to decrypt ACKs sent in 1-RTT packets. + if !h.handshakeConfirmed && !h.hasOutstandingCryptoPackets() { + if h.peerCompletedAddressValidation { + return + } + t := now.Add(h.getScaledPTO(false)) + if h.initialPackets != nil { + return t, protocol.EncryptionInitial + } + return t, protocol.EncryptionHandshake + } + + if h.initialPackets != nil && h.initialPackets.history.HasOutstandingPackets() && + !h.initialPackets.lastAckElicitingPacketTime.IsZero() { + encLevel = protocol.EncryptionInitial + if t := h.initialPackets.lastAckElicitingPacketTime; !t.IsZero() { + pto = t.Add(h.getScaledPTO(false)) + } + } + if h.handshakePackets != nil && h.handshakePackets.history.HasOutstandingPackets() && + !h.handshakePackets.lastAckElicitingPacketTime.IsZero() { + t := h.handshakePackets.lastAckElicitingPacketTime.Add(h.getScaledPTO(false)) + if pto.IsZero() || (!t.IsZero() && t.Before(pto)) { + pto = t + encLevel = protocol.EncryptionHandshake + } + } + if h.handshakeConfirmed && h.appDataPackets.history.HasOutstandingPackets() && + !h.appDataPackets.lastAckElicitingPacketTime.IsZero() { + t := h.appDataPackets.lastAckElicitingPacketTime.Add(h.getScaledPTO(true)) + if pto.IsZero() || (!t.IsZero() && t.Before(pto)) { + pto = t + encLevel = protocol.Encryption1RTT + } + } + return pto, encLevel +} + +func (h *sentPacketHandler) hasOutstandingCryptoPackets() bool { + if h.initialPackets != nil && h.initialPackets.history.HasOutstandingPackets() { + return true + } + if h.handshakePackets != nil && h.handshakePackets.history.HasOutstandingPackets() { + return true + } + return false +} + +func (h *sentPacketHandler) setLossDetectionTimer(now monotime.Time) { + oldAlarm := h.alarm // only needed in case tracing is enabled + newAlarm := h.lossDetectionTime(now) + h.alarm = newAlarm + + hasAlarm := !newAlarm.Time.IsZero() + if !hasAlarm && !oldAlarm.Time.IsZero() { + h.logger.Debugf("Canceling loss detection timer.") + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeCancelled, + }) + } + } + + if h.qlogger != nil && hasAlarm && newAlarm != oldAlarm { + h.qlogger.RecordEvent(qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeSet, + TimerType: newAlarm.TimerType, + EncLevel: newAlarm.EncryptionLevel, + Time: newAlarm.Time.ToTime(), + }) + } +} + +func (h *sentPacketHandler) lossDetectionTime(now monotime.Time) alarmTimer { + // cancel the alarm if no packets are outstanding + if h.peerCompletedAddressValidation && !h.hasOutstandingCryptoPackets() && + !h.appDataPackets.history.HasOutstandingPackets() && !h.appDataPackets.history.HasOutstandingPathProbes() { + return alarmTimer{} + } + + // cancel the alarm if amplification limited + if h.isAmplificationLimited() { + return alarmTimer{} + } + + var pathProbeLossTime monotime.Time + if h.appDataPackets.history.HasOutstandingPathProbes() { + if _, p := h.appDataPackets.history.FirstOutstandingPathProbe(); p != nil { + pathProbeLossTime = p.SendTime.Add(pathProbePacketLossTimeout) + } + } + + // early retransmit timer or time loss detection + lossTime, encLevel := h.getLossTimeAndSpace() + if !lossTime.IsZero() && (pathProbeLossTime.IsZero() || lossTime.Before(pathProbeLossTime)) { + return alarmTimer{ + Time: lossTime, + TimerType: qlog.TimerTypeACK, + EncryptionLevel: encLevel, + } + } + ptoTime, encLevel := h.getPTOTimeAndSpace(now) + if !ptoTime.IsZero() && (pathProbeLossTime.IsZero() || ptoTime.Before(pathProbeLossTime)) { + return alarmTimer{ + Time: ptoTime, + TimerType: qlog.TimerTypePTO, + EncryptionLevel: encLevel, + } + } + if !pathProbeLossTime.IsZero() { + return alarmTimer{ + Time: pathProbeLossTime, + TimerType: qlog.TimerTypePathProbe, + EncryptionLevel: protocol.Encryption1RTT, + } + } + return alarmTimer{} +} + +func (h *sentPacketHandler) detectLostPathProbes(now monotime.Time) { + if !h.appDataPackets.history.HasOutstandingPathProbes() { + return + } + lossTime := now.Add(-pathProbePacketLossTimeout) + // RemovePathProbe cannot be called while iterating. + var lostPathProbes []packetWithPacketNumber + for pn, p := range h.appDataPackets.history.PathProbes() { + if !p.SendTime.After(lossTime) { + lostPathProbes = append(lostPathProbes, packetWithPacketNumber{PacketNumber: pn, packet: p}) + } + } + for _, p := range lostPathProbes { + for _, f := range p.Frames { + f.Handler.OnLost(f.Frame) + } + h.appDataPackets.history.RemovePathProbe(p.PacketNumber) + } +} + +func (h *sentPacketHandler) detectLostPackets(now monotime.Time, encLevel protocol.EncryptionLevel) { + h.lostPacketsInfo = h.lostPacketsInfo[:0] + pnSpace := h.getPacketNumberSpace(encLevel) + pnSpace.lossTime = 0 + + maxRTT := float64(max(h.rttStats.LatestRTT(), h.rttStats.SmoothedRTT())) + lossDelay := time.Duration(timeThreshold * maxRTT) + + // Minimum time of granularity before packets are deemed lost. + lossDelay = max(lossDelay, protocol.TimerGranularity) + + // Packets sent before this time are deemed lost. + lostSendTime := now.Add(-lossDelay) + + cc := h.getCongestionControl() + + priorInFlight := h.bytesInFlight + for pn, p := range pnSpace.history.Packets() { + if pn > pnSpace.largestAcked { + break + } + + var packetLost bool + if !p.SendTime.After(lostSendTime) { + packetLost = true + if !p.isPathProbePacket && p.IsAckEliciting() { + if h.logger.Debug() { + h.logger.Debugf("\tlost packet %d (time threshold)", pn) + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.PacketLost{ + Header: qlog.PacketHeader{ + PacketType: qlog.EncryptionLevelToPacketType(p.EncryptionLevel), + PacketNumber: pn, + }, + Trigger: qlog.PacketLossTimeThreshold, + }) + } + } + } else if pnSpace.history.Difference(pnSpace.largestAcked, pn) >= packetThreshold { + packetLost = true + if !p.isPathProbePacket && p.IsAckEliciting() { + if h.logger.Debug() { + h.logger.Debugf("\tlost packet %d (reordering threshold)", pn) + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.PacketLost{ + Header: qlog.PacketHeader{ + PacketType: qlog.EncryptionLevelToPacketType(p.EncryptionLevel), + PacketNumber: pn, + }, + Trigger: qlog.PacketLossReorderingThreshold, + }) + } + } + } else if pnSpace.lossTime.IsZero() { + // Note: This conditional is only entered once per call + lossTime := p.SendTime.Add(lossDelay) + if h.logger.Debug() { + h.logger.Debugf("\tsetting loss timer for packet %d (%s) to %s (in %s)", pn, encLevel, lossDelay, lossTime) + } + pnSpace.lossTime = lossTime + } + if packetLost { + if encLevel == protocol.Encryption0RTT || encLevel == protocol.Encryption1RTT { + h.lostPackets.Add(pn, p.SendTime) + } + pnSpace.history.DeclareLost(pn) + if !p.isPathProbePacket && p.IsAckEliciting() { + // the bytes in flight need to be reduced no matter if the frames in this packet will be retransmitted + h.removeFromBytesInFlight(p) + h.queueFramesForRetransmission(p) + if !p.IsPathMTUProbePacket { + cc.OnCongestionEvent(pn, p.Length, priorInFlight) + } + h.lostPacketsInfo = append(h.lostPacketsInfo, congestionExt.LostPacketInfo{ + PacketNumber: congestionExt.PacketNumber(pn), + BytesLost: congestionExt.ByteCount(p.Length), + }) + if encLevel == protocol.Encryption1RTT && h.ecnTracker != nil { + h.ecnTracker.LostPacket(pn) + } + } + } + } +} + +func (h *sentPacketHandler) OnLossDetectionTimeout(now monotime.Time) error { + defer h.setLossDetectionTimer(now) + + if h.handshakeConfirmed { + h.detectLostPathProbes(now) + } + + priorInFlight := h.bytesInFlight + earliestLossTime, encLevel := h.getLossTimeAndSpace() + if !earliestLossTime.IsZero() { + if h.logger.Debug() { + h.logger.Debugf("Loss detection alarm fired in loss timer mode. Loss time: %s", earliestLossTime) + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeExpired, + TimerType: qlog.TimerTypeACK, + EncLevel: encLevel, + }) + } + // Early retransmit or time loss detection + h.detectLostPackets(now, encLevel) + + if cex, ok := h.getCongestionControl().(congestion.SendAlgorithmEx); ok && + len(h.lostPacketsInfo) != 0 { + cex.OnCongestionEventEx(priorInFlight, now, nil, h.lostPacketsInfo) + } + return nil + } + + // PTO + // When all outstanding are acknowledged, the alarm is canceled in setLossDetectionTimer. + // However, there's no way to reset the timer in the connection. + // When OnLossDetectionTimeout is called, we therefore need to make sure that there are + // actually packets outstanding. + if h.bytesInFlight == 0 && !h.peerCompletedAddressValidation { + h.ptoCount++ + h.numProbesToSend++ + if h.initialPackets != nil { + h.ptoMode = SendPTOInitial + } else if h.handshakePackets != nil { + h.ptoMode = SendPTOHandshake + } else { + return errors.New("sentPacketHandler BUG: PTO fired, but bytes_in_flight is 0 and Initial and Handshake already dropped") + } + return nil + } + + ptoTime, encLevel := h.getPTOTimeAndSpace(now) + if ptoTime.IsZero() { + return nil + } + ps := h.getPacketNumberSpace(encLevel) + if !ps.history.HasOutstandingPackets() && !ps.history.HasOutstandingPathProbes() && !h.peerCompletedAddressValidation { + return nil + } + h.ptoCount++ + if h.logger.Debug() { + h.logger.Debugf("Loss detection alarm for %s fired in PTO mode. PTO count: %d", encLevel, h.ptoCount) + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeExpired, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + }) + h.qlogger.RecordEvent(qlog.PTOCountUpdated{PTOCount: h.ptoCount}) + } + h.numProbesToSend += 2 + //nolint:exhaustive // We never arm a PTO timer for 0-RTT packets. + switch encLevel { + case protocol.EncryptionInitial: + h.ptoMode = SendPTOInitial + case protocol.EncryptionHandshake: + h.ptoMode = SendPTOHandshake + case protocol.Encryption1RTT: + // skip a packet number in order to elicit an immediate ACK + pn := h.PopPacketNumber(protocol.Encryption1RTT) + h.getPacketNumberSpace(protocol.Encryption1RTT).history.SkippedPacket(pn) + h.ptoMode = SendPTOAppData + default: + return fmt.Errorf("PTO timer in unexpected encryption level: %s", encLevel) + } + return nil +} + +func (h *sentPacketHandler) GetLossDetectionTimeout() monotime.Time { + return h.alarm.Time +} + +func (h *sentPacketHandler) ECNMode(isShortHeaderPacket bool) protocol.ECN { + if !h.enableECN { + return protocol.ECNUnsupported + } + if !isShortHeaderPacket { + return protocol.ECNNon + } + return h.ecnTracker.Mode() +} + +// packetNumberSizingUnit is the packet size the congestion window is expressed +// in when deriving the packet number length. The imitated client coalesces a +// packet into a datagram and then caps the next packet to the room that is +// left, so a nearly full datagram makes the window look like a large number of +// packets and lengthens the packet number for exactly one packet. Too little +// room for a header and it keeps the full size instead. +func (h *sentPacketHandler) packetNumberSizingUnit() protocol.ByteCount { + if h.lastDatagramPadding >= minDatagramRoomForSizing { + return h.lastDatagramPadding + } + return h.maxDatagramSize +} + +func (h *sentPacketHandler) PeekPacketNumber(encLevel protocol.EncryptionLevel) (protocol.PacketNumber, protocol.PacketNumberLen) { + pnSpace := h.getPacketNumberSpace(encLevel) + pn := pnSpace.pns.Peek() + // See section 17.1 of RFC 9000. + if h.shortPacketNumbers { + var cwndPackets protocol.PacketNumber + if size := h.packetNumberSizingUnit(); size > 0 { + cwndPackets = protocol.PacketNumber(h.getCongestionControl().GetCongestionWindow() / size) + } + return pn, protocol.PacketNumberLengthForHeaderChrome(pn, pnSpace.largestAcked, cwndPackets) + } + return pn, protocol.PacketNumberLengthForHeader(pn, pnSpace.largestAcked) +} + +func (h *sentPacketHandler) PopPacketNumber(encLevel protocol.EncryptionLevel) protocol.PacketNumber { + pnSpace := h.getPacketNumberSpace(encLevel) + skipped, pn := pnSpace.pns.Pop() + if skipped { + skippedPN := pn - 1 + pnSpace.history.SkippedPacket(skippedPN) + if h.logger.Debug() { + h.logger.Debugf("Skipping packet number %d", skippedPN) + } + } + return pn +} + +func (h *sentPacketHandler) SendMode(now monotime.Time) SendMode { + numTrackedPackets := h.appDataPackets.history.Len() + if h.initialPackets != nil { + numTrackedPackets += h.initialPackets.history.Len() + } + if h.handshakePackets != nil { + numTrackedPackets += h.handshakePackets.history.Len() + } + + if h.isAmplificationLimited() { + h.logger.Debugf("Amplification window limited. Received %d bytes, already sent out %d bytes", h.bytesReceived, h.bytesSent) + return SendNone + } + // Don't send any packets if we're keeping track of the maximum number of packets. + // Note that since MaxOutstandingSentPackets is smaller than MaxTrackedSentPackets, + // we will stop sending out new data when reaching MaxOutstandingSentPackets, + // but still allow sending of retransmissions and ACKs. + if numTrackedPackets >= protocol.MaxTrackedSentPackets { + if h.logger.Debug() { + h.logger.Debugf("Limited by the number of tracked packets: tracking %d packets, maximum %d", numTrackedPackets, protocol.MaxTrackedSentPackets) + } + return SendNone + } + if h.numProbesToSend > 0 { + return h.ptoMode + } + // Only send ACKs if we're congestion limited. + cc := h.getCongestionControl() + if !cc.CanSend(h.bytesInFlight) { + if h.logger.Debug() { + h.logger.Debugf("Congestion limited: bytes in flight %d, window %d", h.bytesInFlight, cc.GetCongestionWindow()) + } + return SendAck + } + if numTrackedPackets >= protocol.MaxOutstandingSentPackets { + if h.logger.Debug() { + h.logger.Debugf("Max outstanding limited: tracking %d packets, maximum: %d", numTrackedPackets, protocol.MaxOutstandingSentPackets) + } + return SendAck + } + if !cc.HasPacingBudget(now) { + return SendPacingLimited + } + return SendAny +} + +func (h *sentPacketHandler) TimeUntilSend() monotime.Time { + return h.getCongestionControl().TimeUntilSend(h.bytesInFlight) +} + +func (h *sentPacketHandler) SetMaxDatagramSize(s protocol.ByteCount) { + h.maxDatagramSize = s + h.getCongestionControl().SetMaxDatagramSize(s) +} + +func (h *sentPacketHandler) SetLastDatagramPadding(n protocol.ByteCount) { + h.lastDatagramPadding = n +} + +func (h *sentPacketHandler) isAmplificationLimited() bool { + if h.peerAddressValidated { + return false + } + return h.bytesSent >= amplificationFactor*h.bytesReceived +} + +func (h *sentPacketHandler) QueueProbePacket(encLevel protocol.EncryptionLevel) bool { + pnSpace := h.getPacketNumberSpace(encLevel) + pn, p := pnSpace.history.FirstOutstanding() + if p == nil { + return false + } + // TODO: don't declare the packet lost here. + // Keep track of acknowledged frames instead. + // Call DeclareLost before queueFramesForRetransmission, which clears the packet's frames. + pnSpace.history.DeclareLost(pn) + h.removeFromBytesInFlight(p) + h.queueFramesForRetransmission(p) + return true +} + +func (h *sentPacketHandler) queueFramesForRetransmission(p *packet) { + if len(p.Frames) == 0 && len(p.StreamFrames) == 0 { + panic("no frames") + } + for _, f := range p.Frames { + if f.Handler != nil { + f.Handler.OnLost(f.Frame) + } + } + for _, f := range p.StreamFrames { + if f.Handler != nil { + f.Handler.OnLost(f.Frame) + } + } + p.StreamFrames = nil + p.Frames = nil +} + +func (h *sentPacketHandler) ResetForRetry(now monotime.Time) { + h.bytesInFlight = 0 + var firstPacketSendTime monotime.Time + for _, p := range h.initialPackets.history.Packets() { + if firstPacketSendTime.IsZero() { + firstPacketSendTime = p.SendTime + } + if p.IsAckEliciting() { + h.queueFramesForRetransmission(p) + } + } + // All application data packets sent at this point are 0-RTT packets. + // In the case of a Retry, we can assume that the server dropped all of them. + for _, p := range h.appDataPackets.history.Packets() { + if p.IsAckEliciting() { + h.queueFramesForRetransmission(p) + } + } + + // Only use the Retry to estimate the RTT if we didn't send any retransmission for the Initial. + // Otherwise, we don't know which Initial the Retry was sent in response to. + if h.ptoCount == 0 { + // Don't set the RTT to a value lower than 5ms here. + h.rttStats.UpdateRTT(max(minRTTAfterRetry, now.Sub(firstPacketSendTime)), 0) + if h.logger.Debug() { + h.logger.Debugf("\tupdated RTT: %s (σ: %s)", h.rttStats.SmoothedRTT(), h.rttStats.MeanDeviation()) + } + if h.qlogger != nil { + h.qlogMetricsUpdated() + } + } + h.initialPackets = newPacketNumberSpace(h.initialPackets.pns.Peek(), false) + h.appDataPackets = newPacketNumberSpace(h.appDataPackets.pns.Peek(), true) + oldAlarm := h.alarm + h.alarm = alarmTimer{} + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.PTOCountUpdated{PTOCount: 0}) + if !oldAlarm.Time.IsZero() { + h.qlogger.RecordEvent(qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeCancelled, + }) + } + } + h.ptoCount = 0 +} + +func (h *sentPacketHandler) MigratedPath(now monotime.Time, initialMaxDatagramSize protocol.ByteCount) { + h.rttStats.ResetForPathMigration() + for pn, p := range h.appDataPackets.history.Packets() { + h.appDataPackets.history.DeclareLost(pn) + if !p.isPathProbePacket { + h.removeFromBytesInFlight(p) + if p.IsAckEliciting() { + h.queueFramesForRetransmission(p) + } + } + } + for pn := range h.appDataPackets.history.PathProbes() { + h.appDataPackets.history.RemovePathProbe(pn) + } + h.congestion = congestion.NewCubicSender( + congestion.DefaultClock{}, + h.rttStats, + h.connStats, + initialMaxDatagramSize, + true, // use Reno + h.qlogger, + ) + h.setLossDetectionTimer(now) +} + +func (h *sentPacketHandler) getCongestionControl() congestion.SendAlgorithmWithDebugInfos { + h.congestionMutex.RLock() + cc := h.congestion + h.congestionMutex.RUnlock() + return cc +} + +func (h *sentPacketHandler) SetCongestionControl(cc congestionExt.CongestionControl) { + h.congestionMutex.Lock() + cc.SetRTTStatsProvider(h.rttStats) + if ccEx, isEx := cc.(congestionExt.CongestionControlEx); isEx { + h.congestion = &ccAdapterEx{ccEx} + } else { + h.congestion = &ccAdapter{cc} + } + h.congestionMutex.Unlock() +} diff --git a/third_party/quic-go/internal/ackhandler/sent_packet_handler_test.go b/third_party/quic-go/internal/ackhandler/sent_packet_handler_test.go new file mode 100644 index 0000000..abd48aa --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/sent_packet_handler_test.go @@ -0,0 +1,1793 @@ +package ackhandler + +import ( + "encoding/binary" + "fmt" + "math/rand/v2" + "slices" + "testing" + "time" + + "github.com/apernet/quic-go/internal/mocks" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type customFrameHandler struct { + onLost, onAcked func(wire.Frame) +} + +func (h *customFrameHandler) OnLost(f wire.Frame) { + if h.onLost != nil { + h.onLost(f) + } +} + +func (h *customFrameHandler) OnAcked(f wire.Frame) { + if h.onAcked != nil { + h.onAcked(f) + } +} + +type packetTracker struct { + Acked []protocol.PacketNumber + Lost []protocol.PacketNumber +} + +func (t *packetTracker) Reset() { + t.Acked = nil + t.Lost = nil +} + +func (t *packetTracker) NewPingFrame(pn protocol.PacketNumber) Frame { + return Frame{ + Frame: &wire.PingFrame{}, + Handler: &customFrameHandler{ + onAcked: func(wire.Frame) { t.Acked = append(t.Acked, pn) }, + onLost: func(wire.Frame) { t.Lost = append(t.Lost, pn) }, + }, + } +} + +func (h *sentPacketHandler) getBytesInFlight() protocol.ByteCount { + return h.bytesInFlight +} + +func ackRanges(pns ...protocol.PacketNumber) []wire.AckRange { + return appendAckRanges(nil, pns...) +} + +func appendAckRanges(ranges []wire.AckRange, pns ...protocol.PacketNumber) []wire.AckRange { + if len(pns) == 0 { + return ranges + } + slices.Sort(pns) + slices.Reverse(pns) + + start := pns[0] + for i := 1; i < len(pns); i++ { + if pns[i-1]-pns[i] > 1 { + ranges = append(ranges, wire.AckRange{Smallest: pns[i-1], Largest: start}) + start = pns[i] + } + } + return append(ranges, wire.AckRange{Smallest: pns[len(pns)-1], Largest: start}) +} + +func TestAckRanges(t *testing.T) { + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 1}}, ackRanges(1)) + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 2}}, ackRanges(1, 2)) + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 3}}, ackRanges(1, 2, 3)) + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 3}}, ackRanges(3, 2, 1)) + require.Equal(t, []wire.AckRange{{Smallest: 1, Largest: 3}}, ackRanges(1, 3, 2)) + + require.Equal(t, []wire.AckRange{{Smallest: 3, Largest: 3}, {Smallest: 1, Largest: 1}}, ackRanges(1, 3)) + require.Equal(t, []wire.AckRange{{Smallest: 3, Largest: 4}, {Smallest: 1, Largest: 1}}, ackRanges(1, 3, 4)) + require.Equal(t, []wire.AckRange{{Smallest: 5, Largest: 6}, {Smallest: 0, Largest: 2}}, ackRanges(0, 1, 2, 5, 6)) +} + +func TestSentPacketHandlerSendAndAcknowledge(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testSentPacketHandlerSendAndAcknowledge(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testSentPacketHandlerSendAndAcknowledge(t, protocol.EncryptionHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testSentPacketHandlerSendAndAcknowledge(t, protocol.Encryption1RTT) + }) +} + +func testSentPacketHandlerSendAndAcknowledge(t *testing.T, encLevel protocol.EncryptionLevel) { + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + false, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + var packets packetTracker + var pns []protocol.PacketNumber + now := monotime.Now() + for i := range 10 { + e := encLevel + // also send some 0-RTT packets to make sure they're acknowledged in the same packet number space + if encLevel == protocol.Encryption1RTT && i < 5 { + e = protocol.Encryption0RTT + } + pn := sph.PopPacketNumber(e) + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, e, protocol.ECNNon, 1200, false, false) + pns = append(pns, pn) + } + + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[0], pns[1], pns[2], pns[3], pns[4], pns[7], pns[8], pns[9])}, + encLevel, + monotime.Now(), + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[0], pns[1], pns[2], pns[3], pns[4], pns[7], pns[8], pns[9]}, packets.Acked) + + // ACKs that don't acknowledge new packets are ok + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[1], pns[2], pns[3])}, + encLevel, + monotime.Now(), + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[0], pns[1], pns[2], pns[3], pns[4], pns[7], pns[8], pns[9]}, packets.Acked) + + // ACKs that don't acknowledge packets that we didn't send are not ok + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[7], pns[8], pns[9], pns[9]+1)}, + encLevel, + monotime.Now(), + ) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + require.ErrorContains(t, err, "received ACK for an unsent packet") +} + +func TestSentPacketHandlerAcknowledgeSkippedPacket(t *testing.T) { + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + false, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + now := monotime.Now() + lastPN := protocol.InvalidPacketNumber + skippedPN := protocol.InvalidPacketNumber + for { + pn, _ := sph.PeekPacketNumber(protocol.Encryption1RTT) + require.Equal(t, pn, sph.PopPacketNumber(protocol.Encryption1RTT)) + if pn > lastPN+1 { + skippedPN = pn - 1 + } + if pn >= 1e6 { + t.Fatal("expected a skipped packet number") + } + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, protocol.Encryption1RTT, protocol.ECNNon, 1200, false, false) + lastPN = pn + if skippedPN != protocol.InvalidPacketNumber { + break + } + } + + _, err := sph.ReceivedAck(&wire.AckFrame{ + AckRanges: []wire.AckRange{{Smallest: 0, Largest: lastPN}}, + }, protocol.Encryption1RTT, monotime.Now()) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) + require.ErrorContains(t, err, fmt.Sprintf("received an ACK for skipped packet number: %d (1-RTT)", skippedPN)) +} + +func TestSentPacketHandlerRTTAckEliciting(t *testing.T) { + var eventRecorder events.Recorder + + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + false, + false, + nil, + protocol.PerspectiveClient, + false, + &eventRecorder, + utils.DefaultLogger, + ) + + getPacketsInFlight := func() int { + evs := eventRecorder.Events(qlog.MetricsUpdated{}) + return evs[len(evs)-1].(qlog.MetricsUpdated).PacketsInFlight + } + getBytesInFlight := func() int { + evs := eventRecorder.Events(qlog.MetricsUpdated{}) + return evs[len(evs)-1].(qlog.MetricsUpdated).BytesInFlight + } + + sendPacket := func(t *testing.T, ti monotime.Time, size protocol.ByteCount, ackEliciting bool) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + var frames []Frame + if ackEliciting { + frames = []Frame{{Frame: &wire.PingFrame{}}} + } + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, frames, protocol.Encryption1RTT, protocol.ECNNon, size, false, false) + return pn + } + + ackPackets := func(t *testing.T, ti monotime.Time, pns ...protocol.PacketNumber) { + t.Helper() + _, err := sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns...)}, protocol.Encryption1RTT, ti) + require.NoError(t, err) + } + + now := monotime.Now() + pn1 := sendPacket(t, now, 1200, true) + require.Equal(t, 1, getPacketsInFlight()) + require.Equal(t, 1200, getBytesInFlight()) + pn2 := sendPacket(t, now, 1100, false) + // Sending a non-ack-eliciting packet doesn't change bytes or packets in flight. + // Non-ack-eliciting packets are not included in congestion control. + require.Equal(t, 1, getPacketsInFlight()) + require.Equal(t, 1200, getBytesInFlight()) + pn3 := sendPacket(t, now, 1000, true) + require.Equal(t, 2, getPacketsInFlight()) + require.Equal(t, 2200, getBytesInFlight()) + // the RTT is recorded, since the largest acknowledged packet is ack-eliciting + now = now.Add(200 * time.Millisecond) + ackPackets(t, now, pn1, pn2, pn3) + require.Equal(t, 200*time.Millisecond, rttStats.LatestRTT()) + require.Zero(t, getPacketsInFlight()) + require.Zero(t, getBytesInFlight()) + + pn4 := sendPacket(t, now, 1200, false) + // non-ack-eliciting packets don't trigger metrics updates + require.Zero(t, getPacketsInFlight()) + require.Zero(t, getBytesInFlight()) + pn5 := sendPacket(t, now, 500, false) + require.Zero(t, getPacketsInFlight()) + require.Zero(t, getBytesInFlight()) + now = now.Add(500 * time.Millisecond) + // only non-ack-eliciting packets are newly acknowledged, so the RTT is not updated + ackPackets(t, now, pn2, pn3, pn4, pn5) + require.Equal(t, 200*time.Millisecond, rttStats.LatestRTT()) + + pn6 := sendPacket(t, now, 1400, true) + require.Equal(t, 1, getPacketsInFlight()) + require.Equal(t, 1400, getBytesInFlight()) + pn7 := sendPacket(t, now, 1100, false) + // non-ack-eliciting packet doesn't change metrics + require.Equal(t, 1, getPacketsInFlight()) + require.Equal(t, 1400, getBytesInFlight()) + now = now.Add(800 * time.Millisecond) + // largest acknowledged packet is not ack-eliciting, but one new ack-eliciting + // packet was acknowledged, so the RTT is updated + ackPackets(t, now, pn6, pn7) + require.Equal(t, 800*time.Millisecond, rttStats.LatestRTT()) + require.Zero(t, getPacketsInFlight()) + require.Zero(t, getBytesInFlight()) +} + +func TestSentPacketHandlerRTTAcrossPacketNumberSpaces(t *testing.T) { + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + false, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + sendPacket := func(t *testing.T, ti monotime.Time, encLevel protocol.EncryptionLevel) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(encLevel) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, encLevel, protocol.ECNNon, 1200, false, false) + return pn + } + + ackPackets := func(t *testing.T, ti monotime.Time, encLevel protocol.EncryptionLevel, pns ...protocol.PacketNumber) { + t.Helper() + _, err := sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns...)}, encLevel, ti) + require.NoError(t, err) + } + + now := monotime.Now() + initial1 := sendPacket(t, now, protocol.EncryptionInitial) + handshake1 := sendPacket(t, now.Add(time.Second), protocol.EncryptionHandshake) + initial2 := sendPacket(t, now.Add(2*time.Second), protocol.EncryptionInitial) + handshake2 := sendPacket(t, now.Add(2*time.Second), protocol.EncryptionHandshake) + + ackPackets(t, now.Add(3*time.Second), protocol.EncryptionInitial, initial1, initial2) + require.Equal(t, time.Second, rttStats.LatestRTT()) + + // No RTT measurement, since the second initial packet was sent after the first handshake packet. + ackPackets(t, now.Add(4*time.Second), protocol.EncryptionHandshake, handshake1) + require.Equal(t, time.Second, rttStats.LatestRTT()) + + // This causes an RTT measurement, since the second handshake packet was sent last. + ackPackets(t, now.Add(5*time.Second), protocol.EncryptionHandshake, handshake1, handshake2) + require.Equal(t, 3*time.Second, rttStats.LatestRTT()) +} + +func TestSentPacketHandlerRTTAckDelays(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testSentPacketHandlerRTTAckDelays(t, protocol.EncryptionInitial, false) + }) + t.Run("Handshake", func(t *testing.T) { + testSentPacketHandlerRTTAckDelays(t, protocol.EncryptionHandshake, false) + }) + t.Run("1-RTT", func(t *testing.T) { + testSentPacketHandlerRTTAckDelays(t, protocol.Encryption1RTT, true) + }) +} + +func testSentPacketHandlerRTTAckDelays(t *testing.T, encLevel protocol.EncryptionLevel, usesAckDelay bool) { + expectedRTTStats := utils.NewRTTStats() + expectedRTTStats.SetMaxAckDelay(time.Second) + rttStats := utils.NewRTTStats() + rttStats.SetMaxAckDelay(time.Second) + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + false, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + sendPacket := func(t *testing.T, ti monotime.Time) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(encLevel) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, encLevel, protocol.ECNNon, 1200, false, false) + return pn + } + + ackPacket := func(pn protocol.PacketNumber, ti monotime.Time, d time.Duration) { + t.Helper() + _, err := sph.ReceivedAck(&wire.AckFrame{DelayTime: d, AckRanges: ackRanges(pn)}, encLevel, ti) + require.NoError(t, err) + } + + var packets []protocol.PacketNumber + now := monotime.Now() + // send some packets and receive ACKs with 0 ack delay + for range 5 { + packets = append(packets, sendPacket(t, now)) + } + for i := range 5 { + expectedRTTStats.UpdateRTT(time.Duration(i+1)*time.Second, 0) + now = now.Add(time.Second) + ackPacket(packets[i], now, 0) + require.Equal(t, expectedRTTStats.SmoothedRTT(), rttStats.SmoothedRTT()) + require.Equal(t, time.Second, rttStats.MinRTT()) + require.Equal(t, time.Duration(i+1)*time.Second, rttStats.LatestRTT()) + } + packets = packets[:0] + + // send some more packets and receive ACKs with non-zero ack delay + for range 5 { + packets = append(packets, sendPacket(t, now)) + } + expectedRTTStatsNoAckDelay := expectedRTTStats.Clone() + for i := range 5 { + const ackDelay = 500 * time.Millisecond + expectedRTTStats.UpdateRTT(time.Duration(i+1)*time.Second, ackDelay) + expectedRTTStatsNoAckDelay.UpdateRTT(time.Duration(i+1)*time.Second, 0) + now = now.Add(time.Second) + ackPacket(packets[i], now, ackDelay) + if usesAckDelay { + require.Equal(t, expectedRTTStats.SmoothedRTT(), rttStats.SmoothedRTT()) + } else { + require.Equal(t, expectedRTTStatsNoAckDelay.SmoothedRTT(), rttStats.SmoothedRTT()) + } + } + packets = packets[:0] + // make sure that taking ack delay into account actually changes the RTT, + // otherwise the test is not meaningful + require.NotEqual(t, expectedRTTStats.SmoothedRTT(), expectedRTTStatsNoAckDelay.SmoothedRTT()) + + // Send two more packets, and acknowledge them in opposite order. + // This tests that the RTT is updated even if the ACK doesn't increase the largest acked. + packets = append(packets, sendPacket(t, now)) + packets = append(packets, sendPacket(t, now)) + ackPacket(packets[1], now.Add(time.Second), 0) + rtt := rttStats.SmoothedRTT() + ackPacket(packets[0], now.Add(10*time.Second), 0) + require.NotEqual(t, rtt, rttStats.SmoothedRTT()) + + // Send one more packet, and send where the largest acked is acknowledged twice. + pn := sendPacket(t, now) + ackPacket(pn, now.Add(time.Second), 0) + rtt = rttStats.SmoothedRTT() + ackPacket(pn, now.Add(10*time.Second), 0) + require.Equal(t, rtt, rttStats.SmoothedRTT()) +} + +func TestSentPacketHandlerAmplificationLimitServer(t *testing.T) { + t.Run("address validated", func(t *testing.T) { + testSentPacketHandlerAmplificationLimitServer(t, true) + }) + t.Run("address not validated", func(t *testing.T) { + testSentPacketHandlerAmplificationLimitServer(t, false) + }) +} + +func testSentPacketHandlerAmplificationLimitServer(t *testing.T, addressValidated bool) { + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + addressValidated, + false, + nil, + protocol.PerspectiveServer, + false, + nil, + utils.DefaultLogger, + ) + + if addressValidated { + require.Equal(t, SendAny, sph.SendMode(monotime.Now())) + return + } + + // no data received yet, so we can't send any packet yet + require.Equal(t, SendNone, sph.SendMode(monotime.Now())) + require.Zero(t, sph.GetLossDetectionTimeout()) + + // Receive 1000 bytes from the client. + // As long as we haven't sent out 3x the amount of bytes received, we can send out new packets, + // even if we go above the 3x limit by sending the last packet. + sph.ReceivedBytes(1000, monotime.Now()) + for i := range 4 { + require.Equal(t, SendAny, sph.SendMode(monotime.Now())) + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(monotime.Now(), pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, protocol.EncryptionInitial, protocol.ECNNon, 999, false, false) + if i != 3 { + require.NotZero(t, sph.GetLossDetectionTimeout()) + } + } + require.Equal(t, SendNone, sph.SendMode(monotime.Now())) + // no need to set a loss detection timer, as we're blocked by the amplification limit + require.Zero(t, sph.GetLossDetectionTimeout()) + + // receiving more data allows us to send out more packets + sph.ReceivedBytes(1000, monotime.Now()) + require.NotZero(t, sph.GetLossDetectionTimeout()) + for range 3 { + require.Equal(t, SendAny, sph.SendMode(monotime.Now())) + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(monotime.Now(), pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, protocol.EncryptionInitial, protocol.ECNNon, 1000, false, false) + } + require.Equal(t, SendNone, sph.SendMode(monotime.Now())) + require.Zero(t, sph.GetLossDetectionTimeout()) + + // receiving an Initial packet doesn't validate the client's address + sph.ReceivedPacket(protocol.EncryptionInitial, monotime.Now()) + require.Equal(t, SendNone, sph.SendMode(monotime.Now())) + require.Zero(t, sph.GetLossDetectionTimeout()) + + // receiving a Handshake packet validates the client's address + sph.ReceivedPacket(protocol.EncryptionHandshake, monotime.Now()) + require.Equal(t, SendAny, sph.SendMode(monotime.Now())) + require.NotZero(t, sph.GetLossDetectionTimeout()) +} + +func TestSentPacketHandlerAmplificationLimitClient(t *testing.T) { + t.Run("handshake ACK", func(t *testing.T) { + testSentPacketHandlerAmplificationLimitClient(t, false) + }) + + t.Run("drop Handshake without ACK", func(t *testing.T) { + testSentPacketHandlerAmplificationLimitClient(t, true) + }) +} + +func testSentPacketHandlerAmplificationLimitClient(t *testing.T, dropHandshake bool) { + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + require.Equal(t, SendAny, sph.SendMode(monotime.Now())) + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(monotime.Now(), pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, protocol.EncryptionInitial, protocol.ECNNon, 999, false, false) + // it's not surprising that the loss detection timer is set, as this packet might be lost... + require.NotZero(t, sph.GetLossDetectionTimeout()) + // ... but it's still set after receiving an ACK for this packet, + // since we might need to unblock the server's amplification limit + _, err := sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pn)}, protocol.EncryptionInitial, monotime.Now()) + require.NoError(t, err) + require.NotZero(t, sph.GetLossDetectionTimeout()) + require.Equal(t, SendAny, sph.SendMode(monotime.Now())) + + // when the timer expires, we should send a PTO packet + sph.OnLossDetectionTimeout(monotime.Now()) + require.Equal(t, SendPTOInitial, sph.SendMode(monotime.Now())) + require.NotZero(t, sph.GetLossDetectionTimeout()) + + if dropHandshake { + // dropping the handshake packet number space completes the handshake, + // even if no ACK for a handshake packet was received + sph.DropPackets(protocol.EncryptionHandshake, monotime.Now()) + require.Zero(t, sph.GetLossDetectionTimeout()) + return + } + + // once the Initial packet number space is dropped, we need to send a Handshake PTO packet, + // even if we haven't sent any packet in the Handshake packet number space yet + sph.DropPackets(protocol.EncryptionInitial, monotime.Now()) + require.NotZero(t, sph.GetLossDetectionTimeout()) + sph.OnLossDetectionTimeout(monotime.Now()) + require.Equal(t, SendPTOHandshake, sph.SendMode(monotime.Now())) + + // receiving an ACK for a handshake packet shows that the server completed address validation + pn = sph.PopPacketNumber(protocol.EncryptionHandshake) + sph.SentPacket(monotime.Now(), pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, protocol.EncryptionHandshake, protocol.ECNNon, 999, false, false) + require.NotZero(t, sph.GetLossDetectionTimeout()) + _, err = sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pn)}, protocol.EncryptionHandshake, monotime.Now()) + require.NoError(t, err) + require.Zero(t, sph.GetLossDetectionTimeout()) +} + +func TestSentPacketHandlerDelayBasedLossDetection(t *testing.T) { + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveServer, + false, + nil, + utils.DefaultLogger, + ) + + var packets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time, isPathMTUProbePacket bool) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.EncryptionInitial, protocol.ECNNon, 1000, isPathMTUProbePacket, false) + return pn + } + + const rtt = time.Second + now := monotime.Now() + t1 := now.Add(-rtt) + t2 := now.Add(-10 * time.Millisecond) + // Send 3 packets + pn1 := sendPacket(t, t1, false) + pn2 := sendPacket(t, t2, false) + // Also send a Path MTU probe packet. + // We expect the same loss recovery logic to be applied to it. + pn3 := sendPacket(t, t2, true) + pn4 := sendPacket(t, now, false) + + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pn4)}, + protocol.EncryptionInitial, + now.Add(time.Second), + ) + require.NoError(t, err) + // make sure that the RTT is actually 1s + require.Equal(t, rtt, rttStats.SmoothedRTT()) + require.Equal(t, []protocol.PacketNumber{pn4}, packets.Acked) + // only the first packet was lost + require.Equal(t, []protocol.PacketNumber{pn1}, packets.Lost) + // ... but we armed a timer to declare packet 2 lost after 9/8 RTTs + require.Equal(t, t2.Add(time.Second*9/8), sph.GetLossDetectionTimeout()) + + sph.OnLossDetectionTimeout(sph.GetLossDetectionTimeout().Add(-time.Microsecond)) + require.Len(t, packets.Lost, 1) + sph.OnLossDetectionTimeout(sph.GetLossDetectionTimeout()) + require.Equal(t, []protocol.PacketNumber{pn1, pn2, pn3}, packets.Lost) +} + +func TestSentPacketHandlerPacketBasedLossDetection(t *testing.T) { + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveServer, + false, + nil, + utils.DefaultLogger, + ) + + var packets packetTracker + now := monotime.Now() + var pns []protocol.PacketNumber + for range 5 { + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.EncryptionInitial, protocol.ECNNon, 1000, false, false) + pns = append(pns, pn) + } + + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[3])}, + protocol.EncryptionInitial, + now.Add(time.Second), + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[3]}, packets.Acked) + require.Equal(t, []protocol.PacketNumber{pns[0]}, packets.Lost) + + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[4])}, + protocol.EncryptionInitial, + now.Add(time.Second), + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[3], pns[4]}, packets.Acked) + require.Equal(t, []protocol.PacketNumber{pns[0], pns[1]}, packets.Lost) +} + +func TestSentPacketHandlerPTO(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testSentPacketHandlerPTO(t, protocol.EncryptionInitial, SendPTOInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testSentPacketHandlerPTO(t, protocol.EncryptionHandshake, SendPTOHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testSentPacketHandlerPTO(t, protocol.Encryption1RTT, SendPTOAppData) + }) +} + +func testSentPacketHandlerPTO(t *testing.T, encLevel protocol.EncryptionLevel, ptoMode SendMode) { + var packets packetTracker + var eventRecorder events.Recorder + + rttStats := utils.NewRTTStats() + rttStats.SetMaxAckDelay(25 * time.Millisecond) + rttStats.UpdateRTT(500*time.Millisecond, 0) + rttStats.UpdateRTT(1000*time.Millisecond, 0) + rttStats.UpdateRTT(1500*time.Millisecond, 0) + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveServer, + false, + &eventRecorder, + utils.DefaultLogger, + ) + + // in the application-data packet number space, the PTO is only set + if encLevel == protocol.Encryption1RTT { + sph.DropPackets(protocol.EncryptionInitial, monotime.Now()) + sph.DropPackets(protocol.EncryptionHandshake, monotime.Now()) + } + + sendPacket := func(t *testing.T, ti monotime.Time, ackEliciting bool, ptoCount uint) protocol.PacketNumber { + t.Helper() + + pn := sph.PopPacketNumber(encLevel) + if ackEliciting { + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, encLevel, protocol.ECNNon, 1000, false, false) + require.Equal(t, + []qlogwriter.Event{ + qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeSet, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + Time: ti.ToTime().Add(rttStats.PTO(encLevel == protocol.Encryption1RTT) << ptoCount), + }, + }, + eventRecorder.Events(qlog.LossTimerUpdated{}), + ) + eventRecorder.Clear() + } else { + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, nil, encLevel, protocol.ECNNon, 1000, true, false) + require.Empty(t, eventRecorder.Events(qlog.LossTimerUpdated{})) + } + return pn + } + + now := monotime.Now() + sendTimes := []monotime.Time{ + now, + now.Add(100 * time.Millisecond), + now.Add(200 * time.Millisecond), + now.Add(300 * time.Millisecond), + } + var pns []protocol.PacketNumber + // send packet 0, 1, 2, 3 + for i := range 3 { + pns = append(pns, sendPacket(t, sendTimes[i], true, 0)) + } + pns = append(pns, sendPacket(t, sendTimes[3], false, 0)) + + // The PTO includes the max_ack_delay only for the application-data packet number space. + // Make sure that the value is actually different, so this test is meaningful. + require.NotEqual(t, rttStats.PTO(true), rttStats.PTO(false)) + + timeout := sph.GetLossDetectionTimeout() + // the PTO is based on the *last* ack-eliciting packet + require.Equal(t, sendTimes[2].Add(rttStats.PTO(encLevel == protocol.Encryption1RTT)), timeout) + + eventRecorder.Clear() + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, + []qlogwriter.Event{ + qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeExpired, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + }, + qlog.PTOCountUpdated{PTOCount: 1}, + qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeSet, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + Time: sendTimes[2].Add(2 * rttStats.PTO(encLevel == protocol.Encryption1RTT)).ToTime(), + }, + }, + eventRecorder.Events(qlog.PTOCountUpdated{}, qlog.LossTimerUpdated{}), + ) + // PTO timer expiration doesn't declare packets lost + require.Empty(t, packets.Lost) + + now = timeout + require.Equal(t, ptoMode, sph.SendMode(now)) + // queue a probe packet + require.True(t, sph.QueueProbePacket(encLevel)) + require.True(t, sph.QueueProbePacket(encLevel)) + require.True(t, sph.QueueProbePacket(encLevel)) + // there are only two ack-eliciting packets that could be queued + require.False(t, sph.QueueProbePacket(encLevel)) + // Queueing probe packets currently works by declaring them lost. + // We shouldn't do this, but this is how the code is currently written. + require.Equal(t, pns[:3], packets.Lost) + packets.Lost = packets.Lost[:0] + + eventRecorder.Clear() + + // send packet 4 and 6 as probe packets + // 5 doesn't count, since it's not an ack-eliciting packet + sendTimes = append(sendTimes, now.Add(100*time.Millisecond)) + sendTimes = append(sendTimes, now.Add(200*time.Millisecond)) + sendTimes = append(sendTimes, now.Add(300*time.Millisecond)) + require.Equal(t, ptoMode, sph.SendMode(sendTimes[4])) // first probe packet + pns = append(pns, sendPacket(t, sendTimes[4], true, 1)) + require.Equal(t, ptoMode, sph.SendMode(sendTimes[5])) // next probe packet + pns = append(pns, sendPacket(t, sendTimes[5], false, 1)) + require.Equal(t, ptoMode, sph.SendMode(sendTimes[6])) // non-ack-eliciting packet didn't count as a probe packet + pns = append(pns, sendPacket(t, sendTimes[6], true, 1)) + require.Equal(t, SendAny, sph.SendMode(sendTimes[6])) // enough probe packets sent + + timeout = sph.GetLossDetectionTimeout() + // exponential backoff + require.Equal(t, sendTimes[6].Add(2*rttStats.PTO(encLevel == protocol.Encryption1RTT)), timeout) + now = timeout + + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, + []qlogwriter.Event{ + qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeExpired, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + }, + qlog.PTOCountUpdated{PTOCount: 2}, + qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeSet, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + Time: sendTimes[6].Add(4 * rttStats.PTO(encLevel == protocol.Encryption1RTT)).ToTime(), + }, + }, + eventRecorder.Events(qlog.LossTimerUpdated{}, qlog.PTOCountUpdated{}), + ) + eventRecorder.Clear() + // PTO timer expiration doesn't declare packets lost + require.Empty(t, packets.Lost) + + // send packet 7, 8 as probe packets + sendTimes = append(sendTimes, now.Add(100*time.Millisecond)) + sendTimes = append(sendTimes, now.Add(200*time.Millisecond)) + require.Equal(t, ptoMode, sph.SendMode(sendTimes[7])) // first probe packet + pns = append(pns, sendPacket(t, sendTimes[7], true, 2)) + require.Equal(t, ptoMode, sph.SendMode(sendTimes[8])) // next probe packet + pns = append(pns, sendPacket(t, sendTimes[8], true, 2)) + require.Equal(t, SendAny, sph.SendMode(sendTimes[8])) // enough probe packets sent + + timeout = sph.GetLossDetectionTimeout() + + // exponential backoff, again + require.Equal(t, sendTimes[8].Add(4*rttStats.PTO(encLevel == protocol.Encryption1RTT)), timeout) + + eventRecorder.Clear() + + // Receive an ACK for packet 7. + // This now declares packets lost, and leads to arming of a timer for packet 8. + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[7])}, + encLevel, + sendTimes[7].Add(time.Microsecond), + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[7]}, packets.Acked) + require.Equal(t, []protocol.PacketNumber{pns[4], pns[6]}, packets.Lost) + require.Len(t, eventRecorder.Events(qlog.PacketLost{}), 2) + require.Equal(t, + []qlogwriter.Event{ + qlog.PTOCountUpdated{PTOCount: 0}, + }, + eventRecorder.Events(qlog.PTOCountUpdated{})[:1], + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.LossTimerUpdated{ + Type: qlog.LossTimerUpdateTypeSet, + TimerType: qlog.TimerTypePTO, + EncLevel: encLevel, + Time: sendTimes[8].Add(rttStats.PTO(encLevel == protocol.Encryption1RTT)).ToTime(), + }, + }, + eventRecorder.Events(qlog.LossTimerUpdated{}), + ) + require.Contains(t, packets.Acked, pns[7]) + + // The PTO timer is now set for the last remaining packet (8), + // with no exponential backoff. + require.Equal(t, sendTimes[8].Add(rttStats.PTO(encLevel == protocol.Encryption1RTT)), sph.GetLossDetectionTimeout()) + + // Acknowledge the last packet (8). + // This should cancel the loss detection timer since there are no more outstanding packets. + eventRecorder.Clear() + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[8])}, + encLevel, + sendTimes[8].Add(time.Second), + ) + require.NoError(t, err) + require.Contains(t, packets.Acked, pns[8]) + + // The loss detection timer should be cancelled since there are no more outstanding packets. + require.True(t, sph.GetLossDetectionTimeout().IsZero()) + require.Equal(t, + []qlogwriter.Event{ + qlog.LossTimerUpdated{Type: qlog.LossTimerUpdateTypeCancelled}, + }, + eventRecorder.Events(qlog.LossTimerUpdated{}), + ) +} + +func TestSentPacketHandlerPacketNumberSpacesPTO(t *testing.T) { + rttStats := utils.NewRTTStats() + const rtt = time.Second + rttStats.UpdateRTT(rtt, 0) + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveServer, + false, + nil, + utils.DefaultLogger, + ) + + sendPacket := func(t *testing.T, ti monotime.Time, encLevel protocol.EncryptionLevel) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(encLevel) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{{Frame: &wire.PingFrame{}}}, encLevel, protocol.ECNNon, 1000, false, false) + return pn + } + + var initialPNs, handshakePNs [4]protocol.PacketNumber + var initialTimes, handshakeTimes [4]monotime.Time + now := monotime.Now() + initialPNs[0] = sendPacket(t, now, protocol.EncryptionInitial) + initialTimes[0] = now + now = now.Add(100 * time.Millisecond) + handshakePNs[0] = sendPacket(t, now, protocol.EncryptionHandshake) + handshakeTimes[0] = now + now = now.Add(100 * time.Millisecond) + initialPNs[1] = sendPacket(t, now, protocol.EncryptionInitial) + initialTimes[1] = now + now = now.Add(100 * time.Millisecond) + handshakePNs[1] = sendPacket(t, now, protocol.EncryptionHandshake) + handshakeTimes[1] = now + require.Equal(t, protocol.ByteCount(4000), sph.(*sentPacketHandler).getBytesInFlight()) + + // the PTO is the earliest time of the PTO times for both packet number spaces, + // i.e. the 2nd Initial packet sent + timeout := sph.GetLossDetectionTimeout() + require.Equal(t, initialTimes[1].Add(rttStats.PTO(false)), timeout) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOInitial, sph.SendMode(timeout)) + // send a PTO probe packet (Initial) + now = timeout.Add(100 * time.Millisecond) + initialPNs[2] = sendPacket(t, now, protocol.EncryptionInitial) + initialTimes[2] = now + + // the earliest PTO time is now the 2nd Handshake packet + timeout = sph.GetLossDetectionTimeout() + // pto_count is a global property, so there's now an exponential backoff + require.Equal(t, handshakeTimes[1].Add(2*rttStats.PTO(false)), timeout) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOHandshake, sph.SendMode(timeout)) + // send a PTO probe packet (Handshake) + now = timeout.Add(100 * time.Millisecond) + handshakePNs[2] = sendPacket(t, now, protocol.EncryptionHandshake) + handshakeTimes[2] = now + + // the earliest PTO time is now the 3rd Initial packet + timeout = sph.GetLossDetectionTimeout() + require.Equal(t, initialTimes[2].Add(4*rttStats.PTO(false)), timeout) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOInitial, sph.SendMode(timeout)) + + // drop the Initial packet number space + now = timeout.Add(100 * time.Millisecond) + require.Equal(t, protocol.ByteCount(6000), sph.(*sentPacketHandler).getBytesInFlight()) + sph.DropPackets(protocol.EncryptionInitial, now) + require.Equal(t, protocol.ByteCount(3000), sph.(*sentPacketHandler).getBytesInFlight()) + + // Since the Initial packets are gone: + // * the earliest PTO time is now based on the 3rd Handshake packet + // * the PTO count is reset to 0 + timeout = sph.GetLossDetectionTimeout() + require.Equal(t, handshakeTimes[2].Add(rttStats.PTO(false)), timeout) + + // send a 1-RTT packet + now = timeout.Add(100 * time.Millisecond) + sendTime := now + sendPacket(t, now, protocol.Encryption1RTT) + + // until handshake confirmation, the PTO timer is based on the Handshake packet number space + require.Equal(t, timeout, sph.GetLossDetectionTimeout()) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOHandshake, sph.SendMode(now)) + + // Drop Handshake packet number space. + // This confirms the handshake, and the PTO timer is now based on the 1-RTT packet number space. + sph.DropPackets(protocol.EncryptionHandshake, now) + require.Equal(t, sendTime.Add(rttStats.PTO(false)), sph.GetLossDetectionTimeout()) +} + +func TestSentPacketHandler0RTT(t *testing.T) { + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + var appDataPackets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time, encLevel protocol.EncryptionLevel) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(encLevel) + var frames []Frame + if encLevel == protocol.Encryption0RTT || encLevel == protocol.Encryption1RTT { + frames = []Frame{appDataPackets.NewPingFrame(pn)} + } else { + frames = []Frame{{Frame: &wire.PingFrame{}}} + } + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, frames, encLevel, protocol.ECNNon, 1000, false, false) + return pn + } + + now := monotime.Now() + sendPacket(t, now, protocol.Encryption0RTT) + sendPacket(t, now.Add(100*time.Millisecond), protocol.EncryptionHandshake) + sendPacket(t, now.Add(200*time.Millisecond), protocol.Encryption0RTT) + sendPacket(t, now.Add(300*time.Millisecond), protocol.Encryption1RTT) + sendPacket(t, now.Add(400*time.Millisecond), protocol.Encryption1RTT) + require.Equal(t, protocol.ByteCount(5000), sph.(*sentPacketHandler).getBytesInFlight()) + + // The PTO timer is based on the Handshake packet number space, not the 0-RTT packets + timeout := sph.GetLossDetectionTimeout() + require.NotZero(t, timeout) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOHandshake, sph.SendMode(timeout)) + + now = timeout.Add(100 * time.Millisecond) + sph.DropPackets(protocol.Encryption0RTT, now) + require.Equal(t, protocol.ByteCount(3000), sph.(*sentPacketHandler).getBytesInFlight()) + // 0-RTT are discarded, not lost + require.Empty(t, appDataPackets.Lost) +} + +func TestSentPacketHandlerCongestion(t *testing.T) { + mockCtrl := gomock.NewController(t) + cong := mocks.NewMockSendAlgorithmWithDebugInfos(mockCtrl) + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveServer, + false, + nil, + utils.DefaultLogger, + ) + sph.(*sentPacketHandler).congestion = cong + + var packets packetTracker + // Send the first 5 packets: not congestion-limited, not pacing-limited. + // The 2nd packet is a Path MTU Probe packet. + now := monotime.Now() + var bytesInFlight protocol.ByteCount + var pns []protocol.PacketNumber + var sendTimes []monotime.Time + for i := range 5 { + gomock.InOrder( + cong.EXPECT().CanSend(bytesInFlight).Return(true), + cong.EXPECT().HasPacingBudget(now).Return(true), + ) + require.Equal(t, SendAny, sph.SendMode(now)) + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + bytesInFlight += 1000 + cong.EXPECT().OnPacketSent(now, bytesInFlight, pn, protocol.ByteCount(1000), true) + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.EncryptionInitial, protocol.ECNNon, 1000, i == 1, false) + pns = append(pns, pn) + sendTimes = append(sendTimes, now) + now = now.Add(100 * time.Millisecond) + } + + // try to send another packet: not congestion-limited, but pacing-limited + now = now.Add(100 * time.Millisecond) + gomock.InOrder( + cong.EXPECT().CanSend(bytesInFlight).Return(true), + cong.EXPECT().HasPacingBudget(now).Return(false), + ) + require.Equal(t, SendPacingLimited, sph.SendMode(now)) + // the connection would call TimeUntilSend, to find out when a new packet can be sent again + pacingDeadline := now.Add(500 * time.Millisecond) + cong.EXPECT().TimeUntilSend(bytesInFlight).Return(pacingDeadline) + require.Equal(t, pacingDeadline, sph.TimeUntilSend()) + + // try to send another packet, but now we're congestion limited + now = now.Add(100 * time.Millisecond) + cong.EXPECT().CanSend(bytesInFlight).Return(false) + require.Equal(t, SendAck, sph.SendMode(now)) // ACKs are allowed even if congestion limited + + // Receive an ACK for packet 3 and 4 (which declares the 1st and 2nd packet lost). + // However, since the 2nd packet was a Path MTU probe packet, it won't get reported + // to the congestion controller. + ackTime := sendTimes[3].Add(time.Second) + gomock.InOrder( + cong.EXPECT().MaybeExitSlowStart(), + cong.EXPECT().OnCongestionEvent(pns[0], protocol.ByteCount(1000), protocol.ByteCount(5000)), + cong.EXPECT().OnPacketAcked(pns[2], protocol.ByteCount(1000), protocol.ByteCount(5000), ackTime), + cong.EXPECT().OnPacketAcked(pns[3], protocol.ByteCount(1000), protocol.ByteCount(5000), ackTime), + ) + _, err := sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns[2], pns[3])}, protocol.EncryptionInitial, ackTime) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[2], pns[3]}, packets.Acked) + require.Equal(t, []protocol.PacketNumber{pns[0], pns[1]}, packets.Lost) + + // Now receive a (delayed) ACK for the 1st packet. + // Since this packet was already lost, we don't expect any calls to the congestion controller. + _, err = sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns[0])}, protocol.EncryptionInitial, ackTime) + require.NoError(t, err) + + // we should now have a PTO timer armed for the 4th packet + timeout := sph.GetLossDetectionTimeout() + require.NotZero(t, timeout) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOInitial, sph.SendMode(timeout)) + + // send another packet to check that bytes_in_flight was correctly adjusted + now = timeout.Add(100 * time.Millisecond) + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + cong.EXPECT().OnPacketSent(now, protocol.ByteCount(2000), pn, protocol.ByteCount(1000), true) + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.EncryptionInitial, protocol.ECNNon, 1000, false, false) +} + +func TestSentPacketHandlerRetry(t *testing.T) { + t.Run("long RTT measurement", func(t *testing.T) { + testSentPacketHandlerRetry(t, time.Second, time.Second) + }) + + // The estimated RTT should be at least 5ms, even if the RTT measurement is very short. + t.Run("short RTT measurement", func(t *testing.T) { + testSentPacketHandlerRetry(t, minRTTAfterRetry/3, minRTTAfterRetry) + }) +} + +func testSentPacketHandlerRetry(t *testing.T, rtt, expectedRTT time.Duration) { + var initialPackets, appDataPackets packetTracker + + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + start := monotime.Now() + now := start + var initialPNs, appDataPNs []protocol.PacketNumber + // send 2 initial and 2 0-RTT packets + for range 2 { + pn := sph.PopPacketNumber(protocol.EncryptionInitial) + initialPNs = append(initialPNs, pn) + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{initialPackets.NewPingFrame(pn)}, protocol.EncryptionInitial, protocol.ECNNon, 1000, false, false) + now = now.Add(100 * time.Millisecond) + + pn = sph.PopPacketNumber(protocol.Encryption0RTT) + appDataPNs = append(appDataPNs, pn) + sph.SentPacket(now, pn, protocol.InvalidPacketNumber, nil, []Frame{appDataPackets.NewPingFrame(pn)}, protocol.Encryption0RTT, protocol.ECNNon, 1000, false, false) + now = now.Add(100 * time.Millisecond) + } + require.Equal(t, protocol.ByteCount(4000), sph.(*sentPacketHandler).getBytesInFlight()) + require.NotZero(t, sph.GetLossDetectionTimeout()) + + sph.ResetForRetry(start.Add(rtt)) + // receiving a Retry cancels all timers + require.Zero(t, sph.GetLossDetectionTimeout()) + // all packets sent so far are declared lost + require.Equal(t, []protocol.PacketNumber{initialPNs[0], initialPNs[1]}, initialPackets.Lost) + require.Equal(t, []protocol.PacketNumber{appDataPNs[0], appDataPNs[1]}, appDataPackets.Lost) + require.False(t, sph.QueueProbePacket(protocol.EncryptionInitial)) + require.False(t, sph.QueueProbePacket(protocol.Encryption0RTT)) + // the RTT measurement is taken from the first packet sent + require.Equal(t, expectedRTT, rttStats.SmoothedRTT()) + require.Zero(t, sph.(*sentPacketHandler).getBytesInFlight()) + + // packet numbers continue increasing + initialPN, _ := sph.PeekPacketNumber(protocol.EncryptionInitial) + require.Greater(t, initialPN, initialPNs[1]) + appDataPN, _ := sph.PeekPacketNumber(protocol.Encryption0RTT) + require.Greater(t, appDataPN, appDataPNs[1]) +} + +func TestSentPacketHandlerRetryAfterPTO(t *testing.T) { + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + + var packets packetTracker + start := monotime.Now() + now := start + pn1 := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(now, pn1, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn1)}, protocol.EncryptionInitial, protocol.ECNNon, 1000, false, false) + + timeout := sph.GetLossDetectionTimeout() + require.NotZero(t, timeout) + sph.OnLossDetectionTimeout(timeout) + require.Equal(t, SendPTOInitial, sph.SendMode(timeout)) + require.True(t, sph.QueueProbePacket(protocol.EncryptionInitial)) + + // send a retransmission for the first packet + now = timeout.Add(100 * time.Millisecond) + pn2 := sph.PopPacketNumber(protocol.EncryptionInitial) + sph.SentPacket(now, pn2, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn2)}, protocol.EncryptionInitial, protocol.ECNNon, 900, false, false) + + const rtt = time.Second + sph.ResetForRetry(now.Add(rtt)) + + require.Equal(t, []protocol.PacketNumber{pn1, pn2}, packets.Lost) + // no RTT measurement is taken, since the PTO timer fired + require.Equal(t, utils.DefaultInitialRTT, rttStats.SmoothedRTT()) +} + +func TestSentPacketHandlerECN(t *testing.T) { + mockCtrl := gomock.NewController(t) + cong := mocks.NewMockSendAlgorithmWithDebugInfos(mockCtrl) + cong.EXPECT().OnPacketSent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + cong.EXPECT().OnPacketAcked(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + cong.EXPECT().MaybeExitSlowStart().AnyTimes() + ecnHandler := NewMockECNHandler(mockCtrl) + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + sph.(*sentPacketHandler).ecnTracker = ecnHandler + sph.(*sentPacketHandler).congestion = cong + + // ECN marks on non-1-RTT packets are ignored + sph.SentPacket(monotime.Now(), sph.PopPacketNumber(protocol.EncryptionInitial), protocol.InvalidPacketNumber, nil, nil, protocol.EncryptionInitial, protocol.ECT1, 1200, false, false) + sph.SentPacket(monotime.Now(), sph.PopPacketNumber(protocol.EncryptionHandshake), protocol.InvalidPacketNumber, nil, nil, protocol.EncryptionHandshake, protocol.ECT0, 1200, false, false) + sph.SentPacket(monotime.Now(), sph.PopPacketNumber(protocol.Encryption0RTT), protocol.InvalidPacketNumber, nil, nil, protocol.Encryption0RTT, protocol.ECNCE, 1200, false, false) + + var packets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time, ecn protocol.ECN) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + ecnHandler.EXPECT().SentPacket(pn, ecn) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.Encryption1RTT, ecn, 1200, false, false) + return pn + } + + pns := make([]protocol.PacketNumber, 4) + now := monotime.Now() + pns[0] = sendPacket(t, now, protocol.ECT1) + now = now.Add(time.Second) + pns[1] = sendPacket(t, now, protocol.ECT0) + pns[2] = sendPacket(t, now, protocol.ECT0) + pns[3] = sendPacket(t, now, protocol.ECT0) + + // Receive an ACK with a short RTT, such that the first packet is lost. + cong.EXPECT().OnCongestionEvent(gomock.Any(), gomock.Any(), gomock.Any()) + ecnHandler.EXPECT().LostPacket(pns[0]) + ecnHandler.EXPECT().HandleNewlyAcked(gomock.Any(), int64(10), int64(11), int64(12)).DoAndReturn(func(packets []packetWithPacketNumber, _, _, _ int64) bool { + require.Len(t, packets, 2) + require.Equal(t, pns[2], packets[0].PacketNumber) + require.Equal(t, pns[3], packets[1].PacketNumber) + return false + }) + _, err := sph.ReceivedAck( + &wire.AckFrame{ + AckRanges: ackRanges(pns[2], pns[3]), + ECT0: 10, + ECT1: 11, + ECNCE: 12, + }, + protocol.Encryption1RTT, + now.Add(100*time.Millisecond), + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[0]}, packets.Lost) + + // The second packet is still outstanding. + // Receive a (delayed) ACK for it. + // Since the new ECN counts were already reported, ECN marks on this ACK frame are ignored. + now = now.Add(100 * time.Millisecond) + _, err = sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns[1])}, protocol.Encryption1RTT, now) + require.NoError(t, err) + + // Send two more packets, and receive an ACK for the second one. + pns = pns[:2] + pns[0] = sendPacket(t, now, protocol.ECT1) + pns[1] = sendPacket(t, now, protocol.ECT1) + ecnHandler.EXPECT().HandleNewlyAcked(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(packets []packetWithPacketNumber, _, _, _ int64) bool { + require.Len(t, packets, 1) + require.Equal(t, pns[1], packets[0].PacketNumber) + return false + }, + ) + now = now.Add(100 * time.Millisecond) + _, err = sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns[1])}, protocol.Encryption1RTT, now) + require.NoError(t, err) + // Receiving an ACK that covers both packets doesn't cause the ECN marks to be reported, + // since the largest acked didn't increase. + now = now.Add(100 * time.Millisecond) + _, err = sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns[0], pns[1])}, protocol.Encryption1RTT, now) + require.NoError(t, err) + + // Send another packet, and have the ECN handler report congestion. + // This needs to be reported to the congestion controller. + pns = pns[:1] + now = now.Add(time.Second) + pns[0] = sendPacket(t, now, protocol.ECT1) + + gomock.InOrder( + ecnHandler.EXPECT().HandleNewlyAcked(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(true), + cong.EXPECT().OnCongestionEvent(pns[0], protocol.ByteCount(0), gomock.Any()), + ) + _, err = sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(pns[0])}, protocol.Encryption1RTT, now.Add(100*time.Millisecond)) + require.NoError(t, err) +} + +func TestSentPacketHandlerPathProbe(t *testing.T) { + const rtt = 10 * time.Millisecond // RTT of the original path + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + sph.DropPackets(protocol.EncryptionInitial, monotime.Now()) + sph.DropPackets(protocol.EncryptionHandshake, monotime.Now()) + + var packets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time, isPathProbe bool) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.Encryption1RTT, protocol.ECNNon, 1200, false, isPathProbe) + return pn + } + + // send 5 packets: 2 non-probe packets, 1 probe packet, 2 non-probe packets + now := monotime.Now() + var pns [5]protocol.PacketNumber + pns[0] = sendPacket(t, now, false) + now = now.Add(rtt) + pns[1] = sendPacket(t, now, false) + pns[2] = sendPacket(t, now, true) + pathProbeTimeout := now.Add(pathProbePacketLossTimeout) + now = now.Add(rtt) + pns[3] = sendPacket(t, now, false) + now = now.Add(rtt) + pns[4] = sendPacket(t, now, false) + require.Less(t, sph.GetLossDetectionTimeout(), pathProbeTimeout) + + now = now.Add(100 * time.Millisecond) + // make sure that this ACK doesn't declare the path probe packet lost + require.Greater(t, pathProbeTimeout, now) + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[0], pns[3], pns[4])}, + protocol.Encryption1RTT, + now, + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[0], pns[3], pns[4]}, packets.Acked) + // despite having been sent at the same time, the probe packet was not lost + require.Equal(t, []protocol.PacketNumber{pns[1]}, packets.Lost) + + // the timeout is now based on the probe packet + timeout := sph.GetLossDetectionTimeout() + require.Equal(t, pathProbeTimeout, timeout) + require.Zero(t, sph.(*sentPacketHandler).getBytesInFlight()) + pn1 := sendPacket(t, now, false) + pn2 := sendPacket(t, now, false) + require.Equal(t, protocol.ByteCount(2400), sph.(*sentPacketHandler).getBytesInFlight()) + + // send one more non-probe packet + pn := sendPacket(t, now, false) + // the timeout is now based on this packet + require.Less(t, sph.GetLossDetectionTimeout(), pathProbeTimeout) + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[2], pn)}, + protocol.Encryption1RTT, + now, + ) + require.NoError(t, err) + + packets.Lost = packets.Lost[:0] + sph.MigratedPath(now, 1200) + require.Zero(t, sph.(*sentPacketHandler).getBytesInFlight()) + require.Equal(t, utils.DefaultInitialRTT, rttStats.SmoothedRTT()) + require.Equal(t, []protocol.PacketNumber{pn1, pn2}, packets.Lost) +} + +func TestSentPacketHandlerPathProbeAckAndLoss(t *testing.T) { + const rtt = 10 * time.Millisecond // RTT of the original path + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + sph.DropPackets(protocol.EncryptionInitial, monotime.Now()) + sph.DropPackets(protocol.EncryptionHandshake, monotime.Now()) + + var packets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time, isPathProbe bool) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.Encryption1RTT, protocol.ECNNon, 1200, false, isPathProbe) + return pn + } + + now := monotime.Now() + pn1 := sendPacket(t, now, true) + t1 := now + now = now.Add(100 * time.Millisecond) + _ = sendPacket(t, now, true) + t2 := now + now = now.Add(100 * time.Millisecond) + pn3 := sendPacket(t, now, true) + + now = now.Add(100 * time.Millisecond) + require.Equal(t, t1.Add(pathProbePacketLossTimeout), sph.GetLossDetectionTimeout()) + require.NoError(t, sph.OnLossDetectionTimeout(sph.GetLossDetectionTimeout())) + require.Equal(t, []protocol.PacketNumber{pn1}, packets.Lost) + packets.Lost = packets.Lost[:0] + + // receive a delayed ACK for the path probe packet + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pn1, pn3)}, + protocol.Encryption1RTT, + now, + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pn3}, packets.Acked) + require.Empty(t, packets.Lost) + + require.Equal(t, t2.Add(pathProbePacketLossTimeout), sph.GetLossDetectionTimeout()) +} + +// The packet tracking logic is pretty complex. +// We test it with a randomized approach, to make sure that it doesn't panic under any circumstances. +func TestSentPacketHandlerRandomized(t *testing.T) { + seed := uint64(time.Now().UnixNano()) + for i := range 5 { + t.Run(fmt.Sprintf("run %d (seed %d)", i+1, seed), func(t *testing.T) { + testSentPacketHandlerRandomized(t, seed) + }) + seed++ + } +} + +func testSentPacketHandlerRandomized(t *testing.T, seed uint64) { + var b [32]byte + binary.BigEndian.PutUint64(b[:], seed) + r := rand.New(rand.NewChaCha8(b)) + + rttStats := utils.NewRTTStats() + rtt := []time.Duration{10 * time.Millisecond, 100 * time.Millisecond, 1000 * time.Millisecond}[r.IntN(3)] + t.Logf("rtt: %dms", rtt.Milliseconds()) + rttStats.UpdateRTT(rtt, 0) // RTT of the original path + + randDuration := func(min, max time.Duration) time.Duration { + return time.Duration(rand.Int64N(int64(max-min))) + min + } + + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + sph.DropPackets(protocol.EncryptionInitial, monotime.Now()) + sph.DropPackets(protocol.EncryptionHandshake, monotime.Now()) + + var packets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time, isPathProbe bool) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.Encryption1RTT, protocol.ECNNon, 1200, false, isPathProbe) + return pn + } + + now := monotime.Now() + start := now + var pns []protocol.PacketNumber + for range 4 { + isProbe := r.Int()%2 == 0 + pn := sendPacket(t, now, isProbe) + t.Logf("t=%dms: sending packet %d (probe packet: %t)", now.Sub(start).Milliseconds(), pn, isProbe) + pns = append(pns, pn) + now = now.Add(randDuration(0, 500*time.Millisecond)) + if r.Int()%3 == 0 { + sph.OnLossDetectionTimeout(now) + t.Logf("t=%dms: loss detection timeout (lost: %v)", now.Sub(start).Milliseconds(), packets.Lost) + packets.Reset() + now = now.Add(randDuration(0, 500*time.Millisecond)) + } + if r.Int()%3 == 0 { + // acknowledge up to 2 random packet numbers from the pns slice + var ackPns []protocol.PacketNumber + if len(pns) > 0 { + numToAck := min(1+r.IntN(2), len(pns)) + for range numToAck { + ackPns = append(ackPns, pns[r.IntN(len(pns))]) + } + } + if len(ackPns) > 1 { + slices.Sort(ackPns) + ackPns = slices.Compact(ackPns) + } + sph.ReceivedAck(&wire.AckFrame{AckRanges: ackRanges(ackPns...)}, protocol.Encryption1RTT, now) + t.Logf("t=%dms: received ACK for packets %v (acked: %v, lost: %v)", now.Sub(start).Milliseconds(), ackPns, packets.Acked, packets.Lost) + packets.Reset() + now = now.Add(randDuration(0, 500*time.Millisecond)) + } + if r.Int()%10 == 0 { + sph.MigratedPath(now, 1200) + now = now.Add(randDuration(0, 500*time.Millisecond)) + } + } + t.Logf("t=%dms: loss detection timeout (lost: %v)", now.Sub(start).Milliseconds(), packets.Lost) + sph.OnLossDetectionTimeout(now) +} + +func TestSentPacketHandlerSpuriousLoss(t *testing.T) { + const rtt = time.Second + + var eventRecorder events.Recorder + + sph := NewSentPacketHandler( + 0, + 1200, + utils.NewRTTStats(), + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + &eventRecorder, + utils.DefaultLogger, + ) + + var packets packetTracker + sendPacket := func(t *testing.T, ti monotime.Time) protocol.PacketNumber { + t.Helper() + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + sph.SentPacket(ti, pn, protocol.InvalidPacketNumber, nil, []Frame{packets.NewPingFrame(pn)}, protocol.Encryption1RTT, protocol.ECNNon, 1000, false, false) + return pn + } + + start := monotime.Now() + now := start + var pns []protocol.PacketNumber + for range 20 { + pns = append(pns, sendPacket(t, now)) + now = now.Add(10 * time.Millisecond) + } + + now = start.Add(rtt) + _, err := sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[0], pns[6])}, + protocol.Encryption1RTT, + now, + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[0], pns[6]}, packets.Acked) + // pns[4] and pns[5] are not yet declared lost + require.Equal(t, []protocol.PacketNumber{pns[1], pns[2], pns[3]}, packets.Lost) + + packets.Reset() + eventRecorder.Clear() + + const secondAckDelay = 50 * time.Millisecond + + now = now.Add(secondAckDelay) + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[0], pns[1], pns[2], pns[3], pns[4], pns[5], pns[6], pns[12], pns[16])}, + protocol.Encryption1RTT, + now, + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[4], pns[5], pns[12], pns[16]}, packets.Acked) + require.Equal(t, []protocol.PacketNumber{pns[7], pns[8], pns[9], pns[10], pns[11], pns[13]}, packets.Lost) + require.Equal(t, + []qlogwriter.Event{ + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[1], + PacketReordering: 16 - 1, + TimeReordering: rtt + secondAckDelay - 10*time.Millisecond, + }, + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[2], + PacketReordering: 16 - 2, + TimeReordering: rtt + secondAckDelay - 20*time.Millisecond, + }, + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[3], + PacketReordering: 16 - 3, + TimeReordering: rtt + secondAckDelay - 30*time.Millisecond, + }, + }, + eventRecorder.Events(qlog.SpuriousLoss{}), + ) + eventRecorder.Clear() + + now = now.Add(secondAckDelay) + _, err = sph.ReceivedAck( + &wire.AckFrame{AckRanges: ackRanges(pns[0], pns[1], pns[2], pns[3], pns[4], pns[5], pns[6], pns[7], pns[8], pns[9], pns[10], pns[16], pns[17], pns[18])}, + protocol.Encryption1RTT, + now, + ) + require.NoError(t, err) + require.Equal(t, []protocol.PacketNumber{pns[4], pns[5], pns[12], pns[16], pns[17], pns[18]}, packets.Acked) + require.Equal(t, []protocol.PacketNumber{pns[7], pns[8], pns[9], pns[10], pns[11], pns[13], pns[14], pns[15]}, packets.Lost) + + require.Equal(t, + []qlogwriter.Event{ + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[7], + PacketReordering: 18 - 7, + TimeReordering: rtt + 2*secondAckDelay - 70*time.Millisecond, + }, + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[8], + PacketReordering: 18 - 8, + TimeReordering: rtt + 2*secondAckDelay - 80*time.Millisecond, + }, + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[9], + PacketReordering: 18 - 9, + TimeReordering: rtt + 2*secondAckDelay - 90*time.Millisecond, + }, + qlog.SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: pns[10], + PacketReordering: 18 - 10, + TimeReordering: rtt + 2*secondAckDelay - 100*time.Millisecond, + }, + }, + eventRecorder.Events(qlog.SpuriousLoss{}), + ) +} + +func BenchmarkSendAndAcknowledge(b *testing.B) { + b.Run("ack every: 2, in flight: 0", func(b *testing.B) { + benchmarkSendAndAcknowledge(b, 2, 0) + }) + b.Run("ack every: 10, in flight: 100", func(b *testing.B) { + benchmarkSendAndAcknowledge(b, 10, 100) + }) + b.Run("ack every: 100, in flight: 1000", func(b *testing.B) { + benchmarkSendAndAcknowledge(b, 100, 1000) + }) +} + +func benchmarkSendAndAcknowledge(b *testing.B, ackEvery, inFlight int) { + b.ReportAllocs() + + rttStats := utils.NewRTTStats() + sph := NewSentPacketHandler( + 0, + 1200, + rttStats, + &utils.ConnectionStats{}, + true, + false, + nil, + protocol.PerspectiveClient, + false, + nil, + utils.DefaultLogger, + ) + now := monotime.Now() + sph.DropPackets(protocol.EncryptionInitial, now) + sph.DropPackets(protocol.EncryptionHandshake, now) + + streamFrames := []StreamFrame{{Frame: &wire.StreamFrame{}}} + + pns := make([]protocol.PacketNumber, 0, ackEvery+inFlight) + + var counter int + ranges := make([]wire.AckRange, 0, ackEvery) + for b.Loop() { + counter++ + pn := sph.PopPacketNumber(protocol.Encryption1RTT) + sph.SentPacket( + now, + pn, + protocol.InvalidPacketNumber, + streamFrames, + nil, + protocol.Encryption1RTT, + protocol.ECNNon, + 1200, + false, false, + ) + now = now.Add(time.Millisecond) + pns = append(pns, pn) + + if counter > inFlight && counter%ackEvery == 0 { + sph.ReceivedAck( + &wire.AckFrame{AckRanges: appendAckRanges(ranges, pns[:ackEvery]...)}, + protocol.Encryption1RTT, + now, + ) + pns = append(pns[:0], pns[ackEvery:]...) + ranges = ranges[:0] + } + } +} diff --git a/third_party/quic-go/internal/ackhandler/sent_packet_history.go b/third_party/quic-go/internal/ackhandler/sent_packet_history.go new file mode 100644 index 0000000..c0f1b88 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/sent_packet_history.go @@ -0,0 +1,274 @@ +package ackhandler + +import ( + "fmt" + "iter" + "slices" + + "github.com/apernet/quic-go/internal/protocol" +) + +const maxSkippedPackets = 4 + +type sentPacketHistory struct { + packets []*packet + pathProbePackets []packetWithPacketNumber + skippedPackets []protocol.PacketNumber + + numOutstanding int + + firstPacketNumber protocol.PacketNumber + highestPacketNumber protocol.PacketNumber +} + +func newSentPacketHistory(isAppData bool) *sentPacketHistory { + h := &sentPacketHistory{ + highestPacketNumber: protocol.InvalidPacketNumber, + firstPacketNumber: protocol.InvalidPacketNumber, + } + if isAppData { + h.packets = make([]*packet, 0, 32) + h.skippedPackets = make([]protocol.PacketNumber, 0, maxSkippedPackets) + } else { + h.packets = make([]*packet, 0, 6) + } + return h +} + +func (h *sentPacketHistory) checkSequentialPacketNumberUse(pn protocol.PacketNumber) { + if h.highestPacketNumber != protocol.InvalidPacketNumber { + if pn != h.highestPacketNumber+1 { + panic("non-sequential packet number use") + } + } + h.highestPacketNumber = pn + if len(h.packets) == 0 { + h.firstPacketNumber = pn + } +} + +func (h *sentPacketHistory) SkippedPacket(pn protocol.PacketNumber) { + h.checkSequentialPacketNumberUse(pn) + if len(h.packets) > 0 { + h.packets = append(h.packets, nil) + } + if len(h.skippedPackets) == maxSkippedPackets { + h.skippedPackets = slices.Delete(h.skippedPackets, 0, 1) + } + h.skippedPackets = append(h.skippedPackets, pn) +} + +func (h *sentPacketHistory) SentPacket(pn protocol.PacketNumber, p *packet) { + h.checkSequentialPacketNumberUse(pn) + h.packets = append(h.packets, p) + if p.Outstanding() { + h.numOutstanding++ + } +} + +func (h *sentPacketHistory) SentPathProbePacket(pn protocol.PacketNumber, p *packet) { + h.checkSequentialPacketNumberUse(pn) + h.packets = append(h.packets, &packet{isPathProbePacket: true}) + h.pathProbePackets = append(h.pathProbePackets, packetWithPacketNumber{PacketNumber: pn, packet: p}) +} + +func (h *sentPacketHistory) Packets() iter.Seq2[protocol.PacketNumber, *packet] { + return func(yield func(protocol.PacketNumber, *packet) bool) { + // h.firstPacketNumber might be updated in the yield function, + // so we need to save it here. + firstPacketNumber := h.firstPacketNumber + for i, p := range h.packets { + if p == nil { + continue + } + if !yield(firstPacketNumber+protocol.PacketNumber(i), p) { + return + } + } + } +} + +func (h *sentPacketHistory) PathProbes() iter.Seq2[protocol.PacketNumber, *packet] { + return func(yield func(protocol.PacketNumber, *packet) bool) { + for _, p := range h.pathProbePackets { + if !yield(p.PacketNumber, p.packet) { + return + } + } + } +} + +// FirstOutstanding returns the first outstanding packet. +func (h *sentPacketHistory) FirstOutstanding() (protocol.PacketNumber, *packet) { + if !h.HasOutstandingPackets() { + return protocol.InvalidPacketNumber, nil + } + for i, p := range h.packets { + if p != nil && p.Outstanding() { + return h.firstPacketNumber + protocol.PacketNumber(i), p + } + } + return protocol.InvalidPacketNumber, nil +} + +// FirstOutstandingPathProbe returns the first outstanding path probe packet +func (h *sentPacketHistory) FirstOutstandingPathProbe() (protocol.PacketNumber, *packet) { + if len(h.pathProbePackets) == 0 { + return protocol.InvalidPacketNumber, nil + } + return h.pathProbePackets[0].PacketNumber, h.pathProbePackets[0].packet +} + +func (h *sentPacketHistory) SkippedPackets() iter.Seq[protocol.PacketNumber] { + return func(yield func(protocol.PacketNumber) bool) { + for _, p := range h.skippedPackets { + if !yield(p) { + return + } + } + } +} + +func (h *sentPacketHistory) Len() int { + return len(h.packets) +} + +func (h *sentPacketHistory) NumOutstanding() int { + return h.numOutstanding +} + +// Remove removes a packet from the sent packet history. +// It must not be used for skipped packet numbers. +func (h *sentPacketHistory) Remove(pn protocol.PacketNumber) error { + idx, ok := h.getIndex(pn) + if !ok { + return fmt.Errorf("packet %d not found in sent packet history", pn) + } + p := h.packets[idx] + if p.Outstanding() { + h.numOutstanding-- + if h.numOutstanding < 0 { + panic("negative number of outstanding packets") + } + } + h.packets[idx] = nil + // clean up all skipped packets directly before this packet number + var hasPacketBefore bool + for idx > 0 { + idx-- + if h.packets[idx] != nil { + hasPacketBefore = true + break + } + } + if !hasPacketBefore { + h.cleanupStart() + } + if len(h.packets) > 0 && h.packets[0] == nil { + panic("cleanup failed") + } + return nil +} + +// RemovePathProbe removes a path probe packet. +// It scales O(N), but that's ok, since we don't expect to send many path probe packets. +// It is not valid to call this function in IteratePathProbes. +func (h *sentPacketHistory) RemovePathProbe(pn protocol.PacketNumber) *packet { + var packetToDelete *packet + idx := -1 + for i, p := range h.pathProbePackets { + if p.PacketNumber == pn { + packetToDelete = p.packet + idx = i + break + } + } + if idx != -1 { + // don't use slices.Delete, because it zeros the deleted element + copy(h.pathProbePackets[idx:], h.pathProbePackets[idx+1:]) + h.pathProbePackets = h.pathProbePackets[:len(h.pathProbePackets)-1] + } + return packetToDelete +} + +// getIndex gets the index of packet p in the packets slice. +func (h *sentPacketHistory) getIndex(p protocol.PacketNumber) (int, bool) { + if len(h.packets) == 0 { + return 0, false + } + if p < h.firstPacketNumber { + return 0, false + } + index := int(p - h.firstPacketNumber) + if index > len(h.packets)-1 { + return 0, false + } + return index, true +} + +func (h *sentPacketHistory) HasOutstandingPackets() bool { + return h.numOutstanding > 0 +} + +func (h *sentPacketHistory) HasOutstandingPathProbes() bool { + return len(h.pathProbePackets) > 0 +} + +// delete all nil entries at the beginning of the packets slice +func (h *sentPacketHistory) cleanupStart() { + for i, p := range h.packets { + if p != nil { + h.packets = h.packets[i:] + h.firstPacketNumber += protocol.PacketNumber(i) + return + } + } + h.packets = h.packets[:0] + h.firstPacketNumber = protocol.InvalidPacketNumber +} + +func (h *sentPacketHistory) LowestPacketNumber() protocol.PacketNumber { + if len(h.packets) == 0 { + return protocol.InvalidPacketNumber + } + return h.firstPacketNumber +} + +func (h *sentPacketHistory) DeclareLost(pn protocol.PacketNumber) { + idx, ok := h.getIndex(pn) + if !ok { + return + } + p := h.packets[idx] + if p.Outstanding() { + h.numOutstanding-- + if h.numOutstanding < 0 { + panic("negative number of outstanding packets") + } + } + h.packets[idx] = nil + if idx == 0 { + h.cleanupStart() + } +} + +// Difference returns the difference between two packet numbers a and b (a - b), +// taking into account any skipped packet numbers between them. +// +// Note that old skipped packets are garbage collected at some point, +// so this function is not guaranteed to return the correct result after a while. +func (h *sentPacketHistory) Difference(a, b protocol.PacketNumber) protocol.PacketNumber { + diff := a - b + if len(h.skippedPackets) == 0 { + return diff + } + if a < h.skippedPackets[0] || b > h.skippedPackets[len(h.skippedPackets)-1] { + return diff + } + for _, p := range h.skippedPackets { + if p > b && p < a { + diff-- + } + } + return diff +} diff --git a/third_party/quic-go/internal/ackhandler/sent_packet_history_test.go b/third_party/quic-go/internal/ackhandler/sent_packet_history_test.go new file mode 100644 index 0000000..7273843 --- /dev/null +++ b/third_party/quic-go/internal/ackhandler/sent_packet_history_test.go @@ -0,0 +1,343 @@ +package ackhandler + +import ( + "slices" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func ackElicitingPacket() *packet { + return &packet{StreamFrames: []StreamFrame{{Frame: &wire.StreamFrame{StreamID: 1}}}} +} + +func (h *sentPacketHistory) getPacketNumbers() []protocol.PacketNumber { + pns := make([]protocol.PacketNumber, 0, len(h.packets)) + for pn := range h.Packets() { + pns = append(pns, pn) + } + return pns +} + +func TestSentPacketHistoryPacketTracking(t *testing.T) { + t.Run("first packet ack-eliciting", func(t *testing.T) { + testSentPacketHistoryPacketTracking(t, true) + }) + t.Run("first packet non-ack-eliciting", func(t *testing.T) { + testSentPacketHistoryPacketTracking(t, false) + }) +} + +func testSentPacketHistoryPacketTracking(t *testing.T, firstPacketAckEliciting bool) { + hist := newSentPacketHistory(true) + + require.False(t, hist.HasOutstandingPackets()) + if firstPacketAckEliciting { + hist.SentPacket(0, ackElicitingPacket()) + require.True(t, hist.HasOutstandingPackets()) + } else { + hist.SentPacket(0, &packet{}) + require.False(t, hist.HasOutstandingPackets()) + } + hist.SentPacket(1, ackElicitingPacket()) + hist.SentPacket(2, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{0, 1, 2}, hist.getPacketNumbers()) + require.Empty(t, slices.Collect(hist.SkippedPackets())) + require.Equal(t, 3, hist.Len()) + if firstPacketAckEliciting { + require.Equal(t, 3, hist.NumOutstanding()) + } else { + require.Equal(t, 2, hist.NumOutstanding()) + } + + // non-ack-eliciting packets are saved, but don't count as outstanding + hist.SentPacket(3, &packet{}) + hist.SentPacket(4, ackElicitingPacket()) + hist.SentPacket(5, &packet{}) + hist.SentPacket(6, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{0, 1, 2, 3, 4, 5, 6}, hist.getPacketNumbers()) + if firstPacketAckEliciting { + require.Equal(t, 5, hist.NumOutstanding()) + } else { + require.Equal(t, 4, hist.NumOutstanding()) + } + + // handle skipped packet numbers + hist.SkippedPacket(7) + hist.SentPacket(8, ackElicitingPacket()) + hist.SentPacket(9, &packet{}) + hist.SkippedPacket(10) + hist.SentPacket(11, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{0, 1, 2, 3, 4, 5, 6, 8, 9, 11}, hist.getPacketNumbers()) + require.Equal(t, []protocol.PacketNumber{7, 10}, slices.Collect(hist.SkippedPackets())) + require.Equal(t, 12, hist.Len()) + if firstPacketAckEliciting { + require.Equal(t, 7, hist.NumOutstanding()) + } else { + require.Equal(t, 6, hist.NumOutstanding()) + } +} + +func TestSentPacketHistoryNonSequentialPacketNumberUse(t *testing.T) { + hist := newSentPacketHistory(true) + hist.SentPacket(100, ackElicitingPacket()) + require.Panics(t, func() { + hist.SentPacket(102, ackElicitingPacket()) + }) +} + +func TestSentPacketHistoryRemovePackets(t *testing.T) { + hist := newSentPacketHistory(true) + + hist.SentPacket(0, ackElicitingPacket()) + hist.SentPacket(1, ackElicitingPacket()) + hist.SkippedPacket(2) + hist.SkippedPacket(3) + hist.SentPacket(4, ackElicitingPacket()) + hist.SkippedPacket(5) + hist.SentPacket(6, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{0, 1, 4, 6}, hist.getPacketNumbers()) + require.Equal(t, []protocol.PacketNumber{2, 3, 5}, slices.Collect(hist.SkippedPackets())) + + require.NoError(t, hist.Remove(0)) + require.Equal(t, []protocol.PacketNumber{2, 3, 5}, slices.Collect(hist.SkippedPackets())) + require.NoError(t, hist.Remove(1)) + require.Equal(t, []protocol.PacketNumber{4, 6}, hist.getPacketNumbers()) + // skipped packets should be preserved + require.Equal(t, []protocol.PacketNumber{2, 3, 5}, slices.Collect(hist.SkippedPackets())) + + // add one more packet + hist.SentPacket(7, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{4, 6, 7}, hist.getPacketNumbers()) + + // remove last packet and add another + require.NoError(t, hist.Remove(7)) + hist.SentPacket(8, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{4, 6, 8}, hist.getPacketNumbers()) + + // try to remove non-existent packet + err := hist.Remove(9) + require.Error(t, err) + require.EqualError(t, err, "packet 9 not found in sent packet history") + + // only the last 4 skipped packets should be preserved + hist.SkippedPacket(9) + hist.SkippedPacket(10) + hist.SentPacket(11, ackElicitingPacket()) + hist.SkippedPacket(12) + require.Equal(t, []protocol.PacketNumber{5, 9, 10, 12}, slices.Collect(hist.SkippedPackets())) + + // Remove all packets + require.NoError(t, hist.Remove(4)) + require.NoError(t, hist.Remove(6)) + require.NoError(t, hist.Remove(8)) + require.NoError(t, hist.Remove(11)) + require.Empty(t, hist.getPacketNumbers()) + require.Len(t, slices.Collect(hist.SkippedPackets()), 4) + require.False(t, hist.HasOutstandingPackets()) +} + +func TestSentPacketHistoryFirstOutstandingPacket(t *testing.T) { + hist := newSentPacketHistory(true) + + pn, p := hist.FirstOutstanding() + require.Equal(t, protocol.InvalidPacketNumber, pn) + require.Nil(t, p) + + hist.SentPacket(2, ackElicitingPacket()) + hist.SentPacket(3, ackElicitingPacket()) + pn, p = hist.FirstOutstanding() + require.Equal(t, protocol.PacketNumber(2), pn) + require.NotNil(t, p) + + // remove the first packet + hist.Remove(2) + pn, p = hist.FirstOutstanding() + require.Equal(t, protocol.PacketNumber(3), pn) + require.NotNil(t, p) + + // Path MTU packets are not regarded as outstanding + hist = newSentPacketHistory(true) + hist.SentPacket(2, ackElicitingPacket()) + hist.SkippedPacket(3) + p = ackElicitingPacket() + p.IsPathMTUProbePacket = true + hist.SentPacket(4, p) + pn, p = hist.FirstOutstanding() + require.NotNil(t, p) + require.Equal(t, protocol.PacketNumber(2), pn) +} + +func TestSentPacketHistoryIterating(t *testing.T) { + hist := newSentPacketHistory(true) + hist.SkippedPacket(0) + hist.SentPacket(1, ackElicitingPacket()) + hist.SentPacket(2, ackElicitingPacket()) + hist.SentPacket(3, ackElicitingPacket()) + hist.SkippedPacket(4) + hist.SkippedPacket(5) + hist.SentPacket(6, ackElicitingPacket()) + require.Equal(t, []protocol.PacketNumber{0, 4, 5}, slices.Collect(hist.SkippedPackets())) + require.NoError(t, hist.Remove(3)) + + var packets []protocol.PacketNumber + for pn, p := range hist.Packets() { + require.NotNil(t, p) + packets = append(packets, pn) + } + + require.Equal(t, []protocol.PacketNumber{1, 2, 6}, packets) + require.Equal(t, []protocol.PacketNumber{0, 4, 5}, slices.Collect(hist.SkippedPackets())) +} + +func TestSentPacketHistoryDeleteWhileIterating(t *testing.T) { + hist := newSentPacketHistory(true) + hist.SentPacket(0, ackElicitingPacket()) + hist.SentPacket(1, ackElicitingPacket()) + hist.SkippedPacket(2) + hist.SentPacket(3, ackElicitingPacket()) + hist.SkippedPacket(4) + hist.SentPacket(5, ackElicitingPacket()) + + var iterations []protocol.PacketNumber + for pn := range hist.Packets() { + iterations = append(iterations, pn) + switch pn { + case 0: + require.NoError(t, hist.Remove(0)) + case 3: + require.NoError(t, hist.Remove(3)) + } + } + + require.Equal(t, []protocol.PacketNumber{0, 1, 3, 5}, iterations) + require.Equal(t, []protocol.PacketNumber{1, 5}, hist.getPacketNumbers()) + require.Equal(t, []protocol.PacketNumber{2, 4}, slices.Collect(hist.SkippedPackets())) +} + +func TestSentPacketHistoryPathProbes(t *testing.T) { + hist := newSentPacketHistory(true) + hist.SentPacket(0, ackElicitingPacket()) + hist.SentPacket(1, ackElicitingPacket()) + hist.SentPathProbePacket(2, ackElicitingPacket()) + hist.SentPacket(3, ackElicitingPacket()) + hist.SentPacket(4, ackElicitingPacket()) + hist.SentPathProbePacket(5, ackElicitingPacket()) + + getPacketsInHistory := func(t *testing.T) []protocol.PacketNumber { + t.Helper() + var pns []protocol.PacketNumber + for pn, p := range hist.Packets() { + pns = append(pns, pn) + switch pn { + case 2, 5: + require.True(t, p.isPathProbePacket) + default: + require.False(t, p.isPathProbePacket) + } + } + return pns + } + + getPacketsInPathProbeHistory := func(t *testing.T) []protocol.PacketNumber { + t.Helper() + var pns []protocol.PacketNumber + for pn := range hist.PathProbes() { + pns = append(pns, pn) + } + return pns + } + + require.Equal(t, []protocol.PacketNumber{0, 1, 2, 3, 4, 5}, getPacketsInHistory(t)) + require.Equal(t, []protocol.PacketNumber{2, 5}, getPacketsInPathProbeHistory(t)) + + // Removing packets from the regular packet history might happen before the path probe + // is declared lost, as the original path might have a smaller RTT than the path timeout. + // Therefore, the path probe packet is not removed from the path probe history. + require.NoError(t, hist.Remove(0)) + require.NoError(t, hist.Remove(1)) + require.NoError(t, hist.Remove(2)) + require.NoError(t, hist.Remove(3)) + require.Equal(t, []protocol.PacketNumber{4, 5}, getPacketsInHistory(t)) + require.Equal(t, []protocol.PacketNumber{2, 5}, getPacketsInPathProbeHistory(t)) + require.True(t, hist.HasOutstandingPackets()) + require.True(t, hist.HasOutstandingPathProbes()) + pn, p := hist.FirstOutstanding() + require.Equal(t, protocol.PacketNumber(4), pn) + require.NotNil(t, p) + pn, p = hist.FirstOutstandingPathProbe() + require.NotNil(t, p) + require.Equal(t, protocol.PacketNumber(2), pn) + + hist.RemovePathProbe(2) + require.Equal(t, []protocol.PacketNumber{4, 5}, getPacketsInHistory(t)) + require.Equal(t, []protocol.PacketNumber{5}, getPacketsInPathProbeHistory(t)) + require.True(t, hist.HasOutstandingPathProbes()) + pn, p = hist.FirstOutstandingPathProbe() + require.NotNil(t, p) + require.Equal(t, protocol.PacketNumber(5), pn) + + hist.RemovePathProbe(5) + require.Equal(t, []protocol.PacketNumber{4, 5}, getPacketsInHistory(t)) + require.Empty(t, getPacketsInPathProbeHistory(t)) + require.True(t, hist.HasOutstandingPackets()) + require.False(t, hist.HasOutstandingPathProbes()) + pn, p = hist.FirstOutstandingPathProbe() + require.Equal(t, protocol.InvalidPacketNumber, pn) + require.Nil(t, p) + + require.NoError(t, hist.Remove(4)) + require.NoError(t, hist.Remove(5)) + require.Empty(t, getPacketsInHistory(t)) + require.False(t, hist.HasOutstandingPackets()) + pn, p = hist.FirstOutstanding() + require.Equal(t, protocol.InvalidPacketNumber, pn) + require.Nil(t, p) + + // path probe packets are considered outstanding + hist.SentPathProbePacket(6, ackElicitingPacket()) + require.False(t, hist.HasOutstandingPackets()) + require.True(t, hist.HasOutstandingPathProbes()) + pn, p = hist.FirstOutstandingPathProbe() + require.NotNil(t, p) + require.Equal(t, protocol.PacketNumber(6), pn) + + hist.RemovePathProbe(6) + require.False(t, hist.HasOutstandingPackets()) + pn, p = hist.FirstOutstanding() + require.Equal(t, protocol.InvalidPacketNumber, pn) + require.Nil(t, p) + require.False(t, hist.HasOutstandingPathProbes()) + pn, p = hist.FirstOutstandingPathProbe() + require.Equal(t, protocol.InvalidPacketNumber, pn) + require.Nil(t, p) +} + +func TestSentPacketHistoryDifference(t *testing.T) { + hist := newSentPacketHistory(true) + hist.SentPacket(0, &packet{}) + hist.SentPacket(1, ackElicitingPacket()) + hist.SentPacket(2, ackElicitingPacket()) + hist.SentPacket(3, ackElicitingPacket()) + hist.SkippedPacket(4) + hist.SkippedPacket(5) + hist.SentPacket(6, ackElicitingPacket()) + hist.SentPacket(7, &packet{}) + hist.SkippedPacket(8) + hist.SentPacket(9, ackElicitingPacket()) + + require.Zero(t, hist.Difference(1, 1)) + require.Zero(t, hist.Difference(2, 2)) + require.Zero(t, hist.Difference(7, 7)) + + require.Equal(t, protocol.PacketNumber(1), hist.Difference(2, 1)) + require.Equal(t, protocol.PacketNumber(2), hist.Difference(3, 1)) + require.Equal(t, protocol.PacketNumber(3), hist.Difference(4, 1)) + require.Equal(t, protocol.PacketNumber(3), hist.Difference(6, 1)) // 4 and 5 were skipped + require.Equal(t, protocol.PacketNumber(4), hist.Difference(7, 1)) // 4 and 5 were skipped + require.Equal(t, protocol.PacketNumber(3), hist.Difference(7, 2)) // 4 and 5 were skipped + require.Equal(t, protocol.PacketNumber(5), hist.Difference(9, 1)) // 4, 5 and 8 were skipped +} diff --git a/third_party/quic-go/internal/congestion/bandwidth.go b/third_party/quic-go/internal/congestion/bandwidth.go new file mode 100644 index 0000000..107f65e --- /dev/null +++ b/third_party/quic-go/internal/congestion/bandwidth.go @@ -0,0 +1,22 @@ +package congestion + +import ( + "time" + + "github.com/apernet/quic-go/internal/protocol" +) + +// Bandwidth of a connection +type Bandwidth uint64 + +const ( + // BitsPerSecond is 1 bit per second + BitsPerSecond Bandwidth = 1 + // BytesPerSecond is 1 byte per second + BytesPerSecond = 8 * BitsPerSecond +) + +// BandwidthFromDelta calculates the bandwidth from a number of bytes and a time delta +func BandwidthFromDelta(bytes protocol.ByteCount, delta time.Duration) Bandwidth { + return Bandwidth(bytes) * Bandwidth(time.Second) / Bandwidth(delta) * BytesPerSecond +} diff --git a/third_party/quic-go/internal/congestion/bandwidth_test.go b/third_party/quic-go/internal/congestion/bandwidth_test.go new file mode 100644 index 0000000..2545d58 --- /dev/null +++ b/third_party/quic-go/internal/congestion/bandwidth_test.go @@ -0,0 +1,12 @@ +package congestion + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestBandwidthFromDelta(t *testing.T) { + require.Equal(t, 1000*BytesPerSecond, BandwidthFromDelta(1, time.Millisecond)) +} diff --git a/third_party/quic-go/internal/congestion/clock.go b/third_party/quic-go/internal/congestion/clock.go new file mode 100644 index 0000000..8337d6b --- /dev/null +++ b/third_party/quic-go/internal/congestion/clock.go @@ -0,0 +1,20 @@ +package congestion + +import ( + "github.com/apernet/quic-go/internal/monotime" +) + +// A Clock returns the current time +type Clock interface { + Now() monotime.Time +} + +// DefaultClock implements the Clock interface using the Go stdlib clock. +type DefaultClock struct{} + +var _ Clock = DefaultClock{} + +// Now gets the current time +func (DefaultClock) Now() monotime.Time { + return monotime.Now() +} diff --git a/third_party/quic-go/internal/congestion/cubic.go b/third_party/quic-go/internal/congestion/cubic.go new file mode 100644 index 0000000..910d62a --- /dev/null +++ b/third_party/quic-go/internal/congestion/cubic.go @@ -0,0 +1,214 @@ +package congestion + +import ( + "math" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +// This cubic implementation is based on the one found in Chromiums's QUIC +// implementation, in the files net/quic/congestion_control/cubic.{hh,cc}. + +// Constants based on TCP defaults. +// The following constants are in 2^10 fractions of a second instead of ms to +// allow a 10 shift right to divide. + +// 1024*1024^3 (first 1024 is from 0.100^3) +// where 0.100 is 100 ms which is the scaling round trip time. +const ( + cubeScale = 40 + cubeCongestionWindowScale = 410 + cubeFactor = 1 << cubeScale / cubeCongestionWindowScale / maxDatagramSize + // TODO: when re-enabling cubic, make sure to use the actual packet size here + maxDatagramSize = protocol.ByteCount(protocol.InitialPacketSize) +) + +const defaultNumConnections = 1 + +// Default Cubic backoff factor +const beta float32 = 0.7 + +// Additional backoff factor when loss occurs in the concave part of the Cubic +// curve. This additional backoff factor is expected to give up bandwidth to +// new concurrent flows and speed up convergence. +const betaLastMax float32 = 0.85 + +// Cubic implements the cubic algorithm from TCP +type Cubic struct { + clock Clock + + // Number of connections to simulate. + numConnections int + + // Time when this cycle started, after last loss event. + epoch monotime.Time + + // Max congestion window used just before last loss event. + // Note: to improve fairness to other streams an additional back off is + // applied to this value if the new value is below our latest value. + lastMaxCongestionWindow protocol.ByteCount + + // Number of acked bytes since the cycle started (epoch). + ackedBytesCount protocol.ByteCount + + // TCP Reno equivalent congestion window in packets. + estimatedTCPcongestionWindow protocol.ByteCount + + // Origin point of cubic function. + originPointCongestionWindow protocol.ByteCount + + // Time to origin point of cubic function in 2^10 fractions of a second. + timeToOriginPoint uint32 + + // Last congestion window in packets computed by cubic function. + lastTargetCongestionWindow protocol.ByteCount +} + +// NewCubic returns a new Cubic instance +func NewCubic(clock Clock) *Cubic { + c := &Cubic{ + clock: clock, + numConnections: defaultNumConnections, + } + c.Reset() + return c +} + +// Reset is called after a timeout to reset the cubic state +func (c *Cubic) Reset() { + c.epoch = 0 + c.lastMaxCongestionWindow = 0 + c.ackedBytesCount = 0 + c.estimatedTCPcongestionWindow = 0 + c.originPointCongestionWindow = 0 + c.timeToOriginPoint = 0 + c.lastTargetCongestionWindow = 0 +} + +func (c *Cubic) alpha() float32 { + // TCPFriendly alpha is described in Section 3.3 of the CUBIC paper. Note that + // beta here is a cwnd multiplier, and is equal to 1-beta from the paper. + // We derive the equivalent alpha for an N-connection emulation as: + b := c.beta() + return 3 * float32(c.numConnections) * float32(c.numConnections) * (1 - b) / (1 + b) +} + +func (c *Cubic) beta() float32 { + // kNConnectionBeta is the backoff factor after loss for our N-connection + // emulation, which emulates the effective backoff of an ensemble of N + // TCP-Reno connections on a single loss event. The effective multiplier is + // computed as: + return (float32(c.numConnections) - 1 + beta) / float32(c.numConnections) +} + +func (c *Cubic) betaLastMax() float32 { + // betaLastMax is the additional backoff factor after loss for our + // N-connection emulation, which emulates the additional backoff of + // an ensemble of N TCP-Reno connections on a single loss event. The + // effective multiplier is computed as: + return (float32(c.numConnections) - 1 + betaLastMax) / float32(c.numConnections) +} + +// OnApplicationLimited is called on ack arrival when sender is unable to use +// the available congestion window. Resets Cubic state during quiescence. +func (c *Cubic) OnApplicationLimited() { + // When sender is not using the available congestion window, the window does + // not grow. But to be RTT-independent, Cubic assumes that the sender has been + // using the entire window during the time since the beginning of the current + // "epoch" (the end of the last loss recovery period). Since + // application-limited periods break this assumption, we reset the epoch when + // in such a period. This reset effectively freezes congestion window growth + // through application-limited periods and allows Cubic growth to continue + // when the entire window is being used. + c.epoch = 0 +} + +// CongestionWindowAfterPacketLoss computes a new congestion window to use after +// a loss event. Returns the new congestion window in packets. The new +// congestion window is a multiplicative decrease of our current window. +func (c *Cubic) CongestionWindowAfterPacketLoss(currentCongestionWindow protocol.ByteCount) protocol.ByteCount { + if currentCongestionWindow+maxDatagramSize < c.lastMaxCongestionWindow { + // We never reached the old max, so assume we are competing with another + // flow. Use our extra back off factor to allow the other flow to go up. + c.lastMaxCongestionWindow = protocol.ByteCount(c.betaLastMax() * float32(currentCongestionWindow)) + } else { + c.lastMaxCongestionWindow = currentCongestionWindow + } + c.epoch = 0 // Reset time. + return protocol.ByteCount(float32(currentCongestionWindow) * c.beta()) +} + +// CongestionWindowAfterAck computes a new congestion window to use after a received ACK. +// Returns the new congestion window in packets. The new congestion window +// follows a cubic function that depends on the time passed since last +// packet loss. +func (c *Cubic) CongestionWindowAfterAck( + ackedBytes protocol.ByteCount, + currentCongestionWindow protocol.ByteCount, + delayMin time.Duration, + eventTime monotime.Time, +) protocol.ByteCount { + c.ackedBytesCount += ackedBytes + + if c.epoch.IsZero() { + // First ACK after a loss event. + c.epoch = eventTime // Start of epoch. + c.ackedBytesCount = ackedBytes // Reset count. + // Reset estimated_tcp_congestion_window_ to be in sync with cubic. + c.estimatedTCPcongestionWindow = currentCongestionWindow + if c.lastMaxCongestionWindow <= currentCongestionWindow { + c.timeToOriginPoint = 0 + c.originPointCongestionWindow = currentCongestionWindow + } else { + c.timeToOriginPoint = uint32(math.Cbrt(float64(cubeFactor * (c.lastMaxCongestionWindow - currentCongestionWindow)))) + c.originPointCongestionWindow = c.lastMaxCongestionWindow + } + } + + // Change the time unit from microseconds to 2^10 fractions per second. Take + // the round trip time in account. This is done to allow us to use shift as a + // divide operator. + elapsedTime := int64(eventTime.Add(delayMin).Sub(c.epoch)/time.Microsecond) << 10 / (1000 * 1000) + + // Right-shifts of negative, signed numbers have implementation-dependent + // behavior, so force the offset to be positive, as is done in the kernel. + offset := int64(c.timeToOriginPoint) - elapsedTime + if offset < 0 { + offset = -offset + } + + deltaCongestionWindow := protocol.ByteCount(cubeCongestionWindowScale*offset*offset*offset) * maxDatagramSize >> cubeScale + var targetCongestionWindow protocol.ByteCount + if elapsedTime > int64(c.timeToOriginPoint) { + targetCongestionWindow = c.originPointCongestionWindow + deltaCongestionWindow + } else { + targetCongestionWindow = c.originPointCongestionWindow - deltaCongestionWindow + } + // Limit the CWND increase to half the acked bytes. + targetCongestionWindow = min(targetCongestionWindow, currentCongestionWindow+c.ackedBytesCount/2) + + // Increase the window by approximately Alpha * 1 MSS of bytes every + // time we ack an estimated tcp window of bytes. For small + // congestion windows (less than 25), the formula below will + // increase slightly slower than linearly per estimated tcp window + // of bytes. + c.estimatedTCPcongestionWindow += protocol.ByteCount(float32(c.ackedBytesCount) * c.alpha() * float32(maxDatagramSize) / float32(c.estimatedTCPcongestionWindow)) + c.ackedBytesCount = 0 + + // We have a new cubic congestion window. + c.lastTargetCongestionWindow = targetCongestionWindow + + // Compute target congestion_window based on cubic target and estimated TCP + // congestion_window, use highest (fastest). + if targetCongestionWindow < c.estimatedTCPcongestionWindow { + targetCongestionWindow = c.estimatedTCPcongestionWindow + } + return targetCongestionWindow +} + +// SetNumConnections sets the number of emulated connections +func (c *Cubic) SetNumConnections(n int) { + c.numConnections = n +} diff --git a/third_party/quic-go/internal/congestion/cubic_sender.go b/third_party/quic-go/internal/congestion/cubic_sender.go new file mode 100644 index 0000000..b764ac2 --- /dev/null +++ b/third_party/quic-go/internal/congestion/cubic_sender.go @@ -0,0 +1,330 @@ +package congestion + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +const ( + // maxDatagramSize is the default maximum packet size used in the Linux TCP implementation. + // Used in QUIC for congestion window computations in bytes. + initialMaxDatagramSize = protocol.ByteCount(protocol.InitialPacketSize) + maxBurstPackets = 3 + renoBeta = 0.7 // Reno backoff factor. + minCongestionWindowPackets = 2 + initialCongestionWindow = 32 +) + +type cubicSender struct { + hybridSlowStart HybridSlowStart + rttStats *utils.RTTStats + connStats *utils.ConnectionStats + cubic *Cubic + pacer *pacer + clock Clock + + reno bool + + // Track the largest packet that has been sent. + largestSentPacketNumber protocol.PacketNumber + + // Track the largest packet that has been acked. + largestAckedPacketNumber protocol.PacketNumber + + // Track the largest packet number outstanding when a CWND cutback occurs. + largestSentAtLastCutback protocol.PacketNumber + + // Whether the last loss event caused us to exit slowstart. + // Used for stats collection of slowstartPacketsLost + lastCutbackExitedSlowstart bool + + // Congestion window in bytes. + congestionWindow protocol.ByteCount + + // Slow start congestion window in bytes, aka ssthresh. + slowStartThreshold protocol.ByteCount + + // ACK counter for the Reno implementation. + numAckedPackets uint64 + + initialCongestionWindow protocol.ByteCount + initialMaxCongestionWindow protocol.ByteCount + + maxDatagramSize protocol.ByteCount + + lastState qlog.CongestionState + qlogger qlogwriter.Recorder +} + +var ( + _ SendAlgorithm = &cubicSender{} + _ SendAlgorithmWithDebugInfos = &cubicSender{} +) + +// NewCubicSender makes a new cubic sender +func NewCubicSender( + clock Clock, + rttStats *utils.RTTStats, + connStats *utils.ConnectionStats, + initialMaxDatagramSize protocol.ByteCount, + reno bool, + qlogger qlogwriter.Recorder, +) *cubicSender { + return newCubicSender( + clock, + rttStats, + connStats, + reno, + initialMaxDatagramSize, + initialCongestionWindow*initialMaxDatagramSize, + protocol.MaxCongestionWindowPackets*initialMaxDatagramSize, + qlogger, + ) +} + +func newCubicSender( + clock Clock, + rttStats *utils.RTTStats, + connStats *utils.ConnectionStats, + reno bool, + initialMaxDatagramSize, + initialCongestionWindow, + initialMaxCongestionWindow protocol.ByteCount, + qlogger qlogwriter.Recorder, +) *cubicSender { + c := &cubicSender{ + rttStats: rttStats, + connStats: connStats, + largestSentPacketNumber: protocol.InvalidPacketNumber, + largestAckedPacketNumber: protocol.InvalidPacketNumber, + largestSentAtLastCutback: protocol.InvalidPacketNumber, + initialCongestionWindow: initialCongestionWindow, + initialMaxCongestionWindow: initialMaxCongestionWindow, + congestionWindow: initialCongestionWindow, + slowStartThreshold: protocol.MaxByteCount, + cubic: NewCubic(clock), + clock: clock, + reno: reno, + qlogger: qlogger, + maxDatagramSize: initialMaxDatagramSize, + } + c.pacer = newPacer(c.BandwidthEstimate) + if c.qlogger != nil { + c.lastState = qlog.CongestionStateSlowStart + c.qlogger.RecordEvent(qlog.CongestionStateUpdated{ + State: qlog.CongestionStateSlowStart, + }) + } + return c +} + +// TimeUntilSend returns when the next packet should be sent. +func (c *cubicSender) TimeUntilSend(_ protocol.ByteCount) monotime.Time { + return c.pacer.TimeUntilSend() +} + +func (c *cubicSender) HasPacingBudget(now monotime.Time) bool { + return c.pacer.Budget(now) >= c.maxDatagramSize +} + +func (c *cubicSender) maxCongestionWindow() protocol.ByteCount { + return c.maxDatagramSize * protocol.MaxCongestionWindowPackets +} + +func (c *cubicSender) minCongestionWindow() protocol.ByteCount { + return c.maxDatagramSize * minCongestionWindowPackets +} + +func (c *cubicSender) OnPacketSent( + sentTime monotime.Time, + _ protocol.ByteCount, + packetNumber protocol.PacketNumber, + bytes protocol.ByteCount, + isRetransmittable bool, +) { + c.pacer.SentPacket(sentTime, bytes) + if !isRetransmittable { + return + } + c.largestSentPacketNumber = packetNumber + c.hybridSlowStart.OnPacketSent(packetNumber) +} + +func (c *cubicSender) CanSend(bytesInFlight protocol.ByteCount) bool { + return bytesInFlight < c.GetCongestionWindow() +} + +func (c *cubicSender) InRecovery() bool { + return c.largestAckedPacketNumber != protocol.InvalidPacketNumber && c.largestAckedPacketNumber <= c.largestSentAtLastCutback +} + +func (c *cubicSender) InSlowStart() bool { + return c.GetCongestionWindow() < c.slowStartThreshold +} + +func (c *cubicSender) GetCongestionWindow() protocol.ByteCount { + return c.congestionWindow +} + +func (c *cubicSender) MaybeExitSlowStart() { + if c.InSlowStart() && + c.hybridSlowStart.ShouldExitSlowStart(c.rttStats.LatestRTT(), c.rttStats.MinRTT(), c.GetCongestionWindow()/c.maxDatagramSize) { + // exit slow start + c.slowStartThreshold = c.congestionWindow + c.maybeQlogStateChange(qlog.CongestionStateCongestionAvoidance) + } +} + +func (c *cubicSender) OnPacketAcked( + ackedPacketNumber protocol.PacketNumber, + ackedBytes protocol.ByteCount, + priorInFlight protocol.ByteCount, + eventTime monotime.Time, +) { + c.largestAckedPacketNumber = max(ackedPacketNumber, c.largestAckedPacketNumber) + if c.InRecovery() { + return + } + c.maybeIncreaseCwnd(ackedPacketNumber, ackedBytes, priorInFlight, eventTime) + if c.InSlowStart() { + c.hybridSlowStart.OnPacketAcked(ackedPacketNumber) + } +} + +func (c *cubicSender) OnCongestionEvent(packetNumber protocol.PacketNumber, lostBytes, priorInFlight protocol.ByteCount) { + c.connStats.PacketsLost.Add(1) + c.connStats.BytesLost.Add(uint64(lostBytes)) + + // TCP NewReno (RFC6582) says that once a loss occurs, any losses in packets + // already sent should be treated as a single loss event, since it's expected. + if packetNumber <= c.largestSentAtLastCutback { + return + } + c.lastCutbackExitedSlowstart = c.InSlowStart() + c.maybeQlogStateChange(qlog.CongestionStateRecovery) + + if c.reno { + c.congestionWindow = protocol.ByteCount(float64(c.congestionWindow) * renoBeta) + } else { + c.congestionWindow = c.cubic.CongestionWindowAfterPacketLoss(c.congestionWindow) + } + if minCwnd := c.minCongestionWindow(); c.congestionWindow < minCwnd { + c.congestionWindow = minCwnd + } + c.slowStartThreshold = c.congestionWindow + c.largestSentAtLastCutback = c.largestSentPacketNumber + // reset packet count from congestion avoidance mode. We start + // counting again when we're out of recovery. + c.numAckedPackets = 0 +} + +// Called when we receive an ack. Normal TCP tracks how many packets one ack +// represents, but quic has a separate ack for each packet. +func (c *cubicSender) maybeIncreaseCwnd( + _ protocol.PacketNumber, + ackedBytes protocol.ByteCount, + priorInFlight protocol.ByteCount, + eventTime monotime.Time, +) { + // Do not increase the congestion window unless the sender is close to using + // the current window. + if !c.isCwndLimited(priorInFlight) { + c.cubic.OnApplicationLimited() + c.maybeQlogStateChange(qlog.CongestionStateApplicationLimited) + return + } + if c.congestionWindow >= c.maxCongestionWindow() { + return + } + if c.InSlowStart() { + // TCP slow start, exponential growth, increase by one for each ACK. + c.congestionWindow += c.maxDatagramSize + c.maybeQlogStateChange(qlog.CongestionStateSlowStart) + return + } + // Congestion avoidance + c.maybeQlogStateChange(qlog.CongestionStateCongestionAvoidance) + if c.reno { + // Classic Reno congestion avoidance. + c.numAckedPackets++ + if c.numAckedPackets >= uint64(c.congestionWindow/c.maxDatagramSize) { + c.congestionWindow += c.maxDatagramSize + c.numAckedPackets = 0 + } + } else { + c.congestionWindow = min( + c.maxCongestionWindow(), + c.cubic.CongestionWindowAfterAck(ackedBytes, c.congestionWindow, c.rttStats.MinRTT(), eventTime), + ) + } +} + +func (c *cubicSender) isCwndLimited(bytesInFlight protocol.ByteCount) bool { + congestionWindow := c.GetCongestionWindow() + if bytesInFlight >= congestionWindow { + return true + } + availableBytes := congestionWindow - bytesInFlight + slowStartLimited := c.InSlowStart() && bytesInFlight > congestionWindow/2 + return slowStartLimited || availableBytes <= maxBurstPackets*c.maxDatagramSize +} + +// BandwidthEstimate returns the current bandwidth estimate +func (c *cubicSender) BandwidthEstimate() Bandwidth { + srtt := c.rttStats.SmoothedRTT() + if srtt == 0 { + // This should never happen, but if it does, avoid division by zero. + srtt = protocol.TimerGranularity + } + return BandwidthFromDelta(c.GetCongestionWindow(), srtt) +} + +// OnRetransmissionTimeout is called on an retransmission timeout +func (c *cubicSender) OnRetransmissionTimeout(packetsRetransmitted bool) { + c.largestSentAtLastCutback = protocol.InvalidPacketNumber + if !packetsRetransmitted { + return + } + c.hybridSlowStart.Restart() + c.cubic.Reset() + c.slowStartThreshold = c.congestionWindow / 2 + c.congestionWindow = c.minCongestionWindow() +} + +// OnConnectionMigration is called when the connection is migrated (?) +func (c *cubicSender) OnConnectionMigration() { + c.hybridSlowStart.Restart() + c.largestSentPacketNumber = protocol.InvalidPacketNumber + c.largestAckedPacketNumber = protocol.InvalidPacketNumber + c.largestSentAtLastCutback = protocol.InvalidPacketNumber + c.lastCutbackExitedSlowstart = false + c.cubic.Reset() + c.numAckedPackets = 0 + c.congestionWindow = c.initialCongestionWindow + c.slowStartThreshold = c.initialMaxCongestionWindow +} + +func (c *cubicSender) maybeQlogStateChange(new qlog.CongestionState) { + if c.qlogger == nil || new == c.lastState { + return + } + c.qlogger.RecordEvent(qlog.CongestionStateUpdated{State: new}) + c.lastState = new +} + +func (c *cubicSender) SetMaxDatagramSize(s protocol.ByteCount) { + if s < c.maxDatagramSize { + panic(fmt.Sprintf("congestion BUG: decreased max datagram size from %d to %d", c.maxDatagramSize, s)) + } + cwndIsMinCwnd := c.congestionWindow == c.minCongestionWindow() + c.maxDatagramSize = s + if cwndIsMinCwnd { + c.congestionWindow = c.minCongestionWindow() + } + c.pacer.SetMaxDatagramSize(s) +} diff --git a/third_party/quic-go/internal/congestion/cubic_sender_test.go b/third_party/quic-go/internal/congestion/cubic_sender_test.go new file mode 100644 index 0000000..e127941 --- /dev/null +++ b/third_party/quic-go/internal/congestion/cubic_sender_test.go @@ -0,0 +1,594 @@ +package congestion + +import ( + "fmt" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + + "github.com/stretchr/testify/require" +) + +const ( + initialCongestionWindowPackets = 10 + defaultWindowTCP = protocol.ByteCount(initialCongestionWindowPackets) * maxDatagramSize +) + +type mockClock monotime.Time + +func (c *mockClock) Now() monotime.Time { + return monotime.Time(*c) +} + +func (c *mockClock) Advance(d time.Duration) { + *c = mockClock(monotime.Time(*c).Add(d)) +} + +const MaxCongestionWindow = 200 * maxDatagramSize + +type testCubicSender struct { + sender *cubicSender + clock *mockClock + rttStats *utils.RTTStats + bytesInFlight protocol.ByteCount + packetNumber protocol.PacketNumber + ackedPacketNumber protocol.PacketNumber +} + +func newTestCubicSender(cubic bool) *testCubicSender { + var clock mockClock + rttStats := utils.RTTStats{} + return &testCubicSender{ + clock: &clock, + rttStats: &rttStats, + packetNumber: 1, + sender: newCubicSender( + &clock, + &rttStats, + &utils.ConnectionStats{}, + !cubic, + protocol.InitialPacketSize, + initialCongestionWindowPackets*maxDatagramSize, + MaxCongestionWindow, + nil, + ), + } +} + +func (s *testCubicSender) SendAvailableSendWindowLen(packetLength protocol.ByteCount) int { + var packetsSent int + for s.sender.CanSend(s.bytesInFlight) { + s.sender.OnPacketSent(s.clock.Now(), s.bytesInFlight, s.packetNumber, packetLength, true) + s.packetNumber++ + packetsSent++ + s.bytesInFlight += packetLength + } + return packetsSent +} + +func (s *testCubicSender) AckNPackets(n int) { + s.rttStats.UpdateRTT(60*time.Millisecond, 0) + s.sender.MaybeExitSlowStart() + for range n { + s.ackedPacketNumber++ + s.sender.OnPacketAcked(s.ackedPacketNumber, maxDatagramSize, s.bytesInFlight, s.clock.Now()) + } + s.bytesInFlight -= protocol.ByteCount(n) * maxDatagramSize + s.clock.Advance(time.Millisecond) +} + +func (s *testCubicSender) LoseNPacketsLen(n int, packetLength protocol.ByteCount) { + for range n { + s.ackedPacketNumber++ + s.sender.OnCongestionEvent(s.ackedPacketNumber, packetLength, s.bytesInFlight) + } + s.bytesInFlight -= protocol.ByteCount(n) * packetLength +} + +func (s *testCubicSender) LosePacket(number protocol.PacketNumber) { + s.sender.OnCongestionEvent(number, maxDatagramSize, s.bytesInFlight) + s.bytesInFlight -= maxDatagramSize +} + +func (s *testCubicSender) SendAvailableSendWindow() int { + return s.SendAvailableSendWindowLen(maxDatagramSize) +} + +func (s *testCubicSender) LoseNPackets(n int) { + s.LoseNPacketsLen(n, maxDatagramSize) +} + +func TestCubicSenderStartup(t *testing.T) { + sender := newTestCubicSender(false) + + // At startup make sure we are at the default. + require.Equal(t, defaultWindowTCP, sender.sender.GetCongestionWindow()) + + // Make sure we can send. + require.Zero(t, sender.sender.TimeUntilSend(0)) + require.True(t, sender.sender.CanSend(sender.bytesInFlight)) + + // And that window is un-affected. + require.Equal(t, defaultWindowTCP, sender.sender.GetCongestionWindow()) + + // Fill the send window with data, then verify that we can't send. + sender.SendAvailableSendWindow() + require.False(t, sender.sender.CanSend(sender.bytesInFlight)) +} + +func TestCubicSenderPacing(t *testing.T) { + sender := newTestCubicSender(false) + + // Set up RTT and advance clock + sender.rttStats.UpdateRTT(10*time.Millisecond, 0) + sender.clock.Advance(time.Hour) + + // Fill the send window with data, then verify that we can't send. + sender.SendAvailableSendWindow() + sender.AckNPackets(1) + + // Check that we can't send immediately due to pacing + delay := sender.sender.TimeUntilSend(sender.bytesInFlight) + require.NotZero(t, delay) + require.Less(t, delay.Sub(monotime.Time(*sender.clock)), time.Hour) +} + +func TestCubicSenderApplicationLimitedSlowStart(t *testing.T) { + sender := newTestCubicSender(false) + + // At startup make sure we can send. + require.True(t, sender.sender.CanSend(0)) + require.Zero(t, sender.sender.TimeUntilSend(0)) + + // Send exactly 10 packets and ensure the CWND ends at 14 packets. + const numberOfAcks = 5 + sender.SendAvailableSendWindow() + for range numberOfAcks { + sender.AckNPackets(2) + } + + bytesToSend := sender.sender.GetCongestionWindow() + // It's expected 2 acks will arrive when the bytes_in_flight are greater than + // half the CWND. + require.Equal(t, defaultWindowTCP+maxDatagramSize*2*2, bytesToSend) +} + +func TestCubicSenderExponentialSlowStart(t *testing.T) { + sender := newTestCubicSender(false) + + // At startup make sure we can send. + require.True(t, sender.sender.CanSend(0)) + require.Zero(t, sender.sender.TimeUntilSend(0)) + + const numberOfAcks = 20 + for range numberOfAcks { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + } + + cwnd := sender.sender.GetCongestionWindow() + require.Equal(t, defaultWindowTCP+maxDatagramSize*2*numberOfAcks, cwnd) + require.Equal(t, BandwidthFromDelta(cwnd, sender.rttStats.SmoothedRTT()), sender.sender.BandwidthEstimate()) +} + +func TestCubicSenderSlowStartPacketLoss(t *testing.T) { + sender := newTestCubicSender(false) + + const numberOfAcks = 10 + for range numberOfAcks { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + } + sender.SendAvailableSendWindow() + expectedSendWindow := defaultWindowTCP + (maxDatagramSize * 2 * numberOfAcks) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Lose a packet to exit slow start. + sender.LoseNPackets(1) + packetsInRecoveryWindow := expectedSendWindow / maxDatagramSize + + // We should now have fallen out of slow start with a reduced window. + expectedSendWindow = protocol.ByteCount(float32(expectedSendWindow) * renoBeta) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Recovery phase. We need to ack every packet in the recovery window before + // we exit recovery. + numberOfPacketsInWindow := expectedSendWindow / maxDatagramSize + sender.AckNPackets(int(packetsInRecoveryWindow)) + sender.SendAvailableSendWindow() + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // We need to ack an entire window before we increase CWND by 1. + fmt.Println(numberOfPacketsInWindow) + sender.AckNPackets(int(numberOfPacketsInWindow) - 2) + sender.SendAvailableSendWindow() + fmt.Println(sender.clock.Now()) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Next ack should increase cwnd by 1. + sender.AckNPackets(1) + expectedSendWindow += maxDatagramSize + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Now RTO and ensure slow start gets reset. + require.True(t, sender.sender.hybridSlowStart.Started()) + sender.sender.OnRetransmissionTimeout(true) + require.False(t, sender.sender.hybridSlowStart.Started()) +} + +func TestCubicSenderSlowStartPacketLossPRR(t *testing.T) { + sender := newTestCubicSender(false) + + // Test based on the first example in RFC6937. + // Ack 10 packets in 5 acks to raise the CWND to 20, as in the example. + const numberOfAcks = 5 + for range numberOfAcks { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + } + sender.SendAvailableSendWindow() + expectedSendWindow := defaultWindowTCP + (maxDatagramSize * 2 * numberOfAcks) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + sender.LoseNPackets(1) + + // We should now have fallen out of slow start with a reduced window. + sendWindowBeforeLoss := expectedSendWindow + expectedSendWindow = protocol.ByteCount(float32(expectedSendWindow) * renoBeta) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Testing TCP proportional rate reduction. + // We should send packets paced over the received acks for the remaining + // outstanding packets. The number of packets before we exit recovery is the + // original CWND minus the packet that has been lost and the one which + // triggered the loss. + remainingPacketsInRecovery := sendWindowBeforeLoss/maxDatagramSize - 2 + + for range remainingPacketsInRecovery { + sender.AckNPackets(1) + sender.SendAvailableSendWindow() + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + } + + // We need to ack another window before we increase CWND by 1. + numberOfPacketsInWindow := expectedSendWindow / maxDatagramSize + for range numberOfPacketsInWindow { + sender.AckNPackets(1) + require.Equal(t, 1, sender.SendAvailableSendWindow()) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + } + + sender.AckNPackets(1) + expectedSendWindow += maxDatagramSize + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) +} + +func TestCubicSenderSlowStartBurstPacketLossPRR(t *testing.T) { + sender := newTestCubicSender(false) + + // Test based on the second example in RFC6937, though we also implement + // forward acknowledgements, so the first two incoming acks will trigger + // PRR immediately. + // Ack 20 packets in 10 acks to raise the CWND to 30. + const numberOfAcks = 10 + for range numberOfAcks { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + } + sender.SendAvailableSendWindow() + expectedSendWindow := defaultWindowTCP + (maxDatagramSize * 2 * numberOfAcks) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Lose one more than the congestion window reduction, so that after loss, + // bytes_in_flight is lesser than the congestion window. + sendWindowAfterLoss := protocol.ByteCount(renoBeta * float32(expectedSendWindow)) + numPacketsToLose := (expectedSendWindow-sendWindowAfterLoss)/maxDatagramSize + 1 + sender.LoseNPackets(int(numPacketsToLose)) + // Immediately after the loss, ensure at least one packet can be sent. + // Losses without subsequent acks can occur with timer based loss detection. + require.True(t, sender.sender.CanSend(sender.bytesInFlight)) + sender.AckNPackets(1) + + // We should now have fallen out of slow start with a reduced window. + expectedSendWindow = protocol.ByteCount(float32(expectedSendWindow) * renoBeta) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Only 2 packets should be allowed to be sent, per PRR-SSRB + require.Equal(t, 2, sender.SendAvailableSendWindow()) + + // Ack the next packet, which triggers another loss. + sender.LoseNPackets(1) + sender.AckNPackets(1) + + // Send 2 packets to simulate PRR-SSRB. + require.Equal(t, 2, sender.SendAvailableSendWindow()) + + // Ack the next packet, which triggers another loss. + sender.LoseNPackets(1) + sender.AckNPackets(1) + + // Send 2 packets to simulate PRR-SSRB. + require.Equal(t, 2, sender.SendAvailableSendWindow()) + + // Exit recovery and return to sending at the new rate. + for range numberOfAcks { + sender.AckNPackets(1) + require.Equal(t, 1, sender.SendAvailableSendWindow()) + } +} + +func TestCubicSenderRTOCongestionWindow(t *testing.T) { + sender := newTestCubicSender(false) + + require.Equal(t, defaultWindowTCP, sender.sender.GetCongestionWindow()) + require.Equal(t, protocol.MaxByteCount, sender.sender.slowStartThreshold) + + // Expect the window to decrease to the minimum once the RTO fires + // and slow start threshold to be set to 1/2 of the CWND. + sender.sender.OnRetransmissionTimeout(true) + require.Equal(t, 2*maxDatagramSize, sender.sender.GetCongestionWindow()) + require.Equal(t, 5*maxDatagramSize, sender.sender.slowStartThreshold) +} + +func TestCubicSenderTCPCubicResetEpochOnQuiescence(t *testing.T) { + sender := newTestCubicSender(true) + + const maxCongestionWindow = 50 + const maxCongestionWindowBytes = maxCongestionWindow * maxDatagramSize + + numSent := sender.SendAvailableSendWindow() + + // Make sure we fall out of slow start. + savedCwnd := sender.sender.GetCongestionWindow() + sender.LoseNPackets(1) + require.Greater(t, savedCwnd, sender.sender.GetCongestionWindow()) + + // Ack the rest of the outstanding packets to get out of recovery. + for i := 1; i < numSent; i++ { + sender.AckNPackets(1) + } + require.Zero(t, sender.bytesInFlight) + + // Send a new window of data and ack all; cubic growth should occur. + savedCwnd = sender.sender.GetCongestionWindow() + numSent = sender.SendAvailableSendWindow() + for range numSent { + sender.AckNPackets(1) + } + require.Less(t, savedCwnd, sender.sender.GetCongestionWindow()) + require.Greater(t, maxCongestionWindowBytes, sender.sender.GetCongestionWindow()) + require.Zero(t, sender.bytesInFlight) + + // Quiescent time of 100 seconds + sender.clock.Advance(100 * time.Second) + + // Send new window of data and ack one packet. Cubic epoch should have + // been reset; ensure cwnd increase is not dramatic. + savedCwnd = sender.sender.GetCongestionWindow() + sender.SendAvailableSendWindow() + sender.AckNPackets(1) + require.InDelta(t, float64(savedCwnd), float64(sender.sender.GetCongestionWindow()), float64(maxDatagramSize)) + require.Greater(t, maxCongestionWindowBytes, sender.sender.GetCongestionWindow()) +} + +func TestCubicSenderMultipleLossesInOneWindow(t *testing.T) { + sender := newTestCubicSender(false) + + sender.SendAvailableSendWindow() + initialWindow := sender.sender.GetCongestionWindow() + sender.LosePacket(sender.ackedPacketNumber + 1) + postLossWindow := sender.sender.GetCongestionWindow() + require.True(t, initialWindow > postLossWindow) + sender.LosePacket(sender.ackedPacketNumber + 3) + require.Equal(t, postLossWindow, sender.sender.GetCongestionWindow()) + sender.LosePacket(sender.packetNumber - 1) + require.Equal(t, postLossWindow, sender.sender.GetCongestionWindow()) + + // Lose a later packet and ensure the window decreases. + sender.LosePacket(sender.packetNumber) + require.True(t, postLossWindow > sender.sender.GetCongestionWindow()) +} + +func TestCubicSender1ConnectionCongestionAvoidanceAtEndOfRecovery(t *testing.T) { + sender := newTestCubicSender(false) + + // Ack 10 packets in 5 acks to raise the CWND to 20. + const numberOfAcks = 5 + for range numberOfAcks { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + } + sender.SendAvailableSendWindow() + expectedSendWindow := defaultWindowTCP + (maxDatagramSize * 2 * numberOfAcks) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + sender.LoseNPackets(1) + + // We should now have fallen out of slow start with a reduced window. + expectedSendWindow = protocol.ByteCount(float32(expectedSendWindow) * renoBeta) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // No congestion window growth should occur in recovery phase, i.e., until the + // currently outstanding 20 packets are acked. + for range 10 { + // Send our full send window. + sender.SendAvailableSendWindow() + require.True(t, sender.sender.InRecovery()) + sender.AckNPackets(2) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + } + require.False(t, sender.sender.InRecovery()) + + // Out of recovery now. Congestion window should not grow during RTT. + for i := protocol.ByteCount(0); i < expectedSendWindow/maxDatagramSize-2; i += 2 { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + } + + // Next ack should cause congestion window to grow by 1MSS. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + expectedSendWindow += maxDatagramSize + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) +} + +func TestCubicSenderNoPRR(t *testing.T) { + sender := newTestCubicSender(false) + + sender.SendAvailableSendWindow() + sender.LoseNPackets(9) + sender.AckNPackets(1) + + require.Equal(t, protocol.ByteCount(renoBeta*float32(defaultWindowTCP)), sender.sender.GetCongestionWindow()) + windowInPackets := int(renoBeta * float32(defaultWindowTCP) / float32(maxDatagramSize)) + numSent := sender.SendAvailableSendWindow() + require.Equal(t, windowInPackets, numSent) +} + +func TestCubicSenderResetAfterConnectionMigration(t *testing.T) { + sender := newTestCubicSender(false) + + require.Equal(t, defaultWindowTCP, sender.sender.GetCongestionWindow()) + require.Equal(t, protocol.MaxByteCount, sender.sender.slowStartThreshold) + + // Starts with slow start. + const numberOfAcks = 10 + for range numberOfAcks { + // Send our full send window. + sender.SendAvailableSendWindow() + sender.AckNPackets(2) + } + sender.SendAvailableSendWindow() + expectedSendWindow := defaultWindowTCP + (maxDatagramSize * 2 * numberOfAcks) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + + // Loses a packet to exit slow start. + sender.LoseNPackets(1) + + // We should now have fallen out of slow start with a reduced window. Slow + // start threshold is also updated. + expectedSendWindow = protocol.ByteCount(float32(expectedSendWindow) * renoBeta) + require.Equal(t, expectedSendWindow, sender.sender.GetCongestionWindow()) + require.Equal(t, expectedSendWindow, sender.sender.slowStartThreshold) + + // Resets cwnd and slow start threshold on connection migrations. + sender.sender.OnConnectionMigration() + require.Equal(t, defaultWindowTCP, sender.sender.GetCongestionWindow()) + require.Equal(t, MaxCongestionWindow, sender.sender.slowStartThreshold) + require.False(t, sender.sender.hybridSlowStart.Started()) +} + +func TestCubicSenderSlowStartsUpToMaximumCongestionWindow(t *testing.T) { + var clock mockClock + rttStats := utils.RTTStats{} + const initialMaxCongestionWindow = protocol.MaxCongestionWindowPackets * initialMaxDatagramSize + sender := newCubicSender( + &clock, + &rttStats, + &utils.ConnectionStats{}, + true, + protocol.InitialPacketSize, + initialCongestionWindowPackets*maxDatagramSize, + initialMaxCongestionWindow, + nil, + ) + + for i := 1; i < protocol.MaxCongestionWindowPackets; i++ { + sender.MaybeExitSlowStart() + sender.OnPacketAcked(protocol.PacketNumber(i), 1350, sender.GetCongestionWindow(), clock.Now()) + } + require.Equal(t, initialMaxCongestionWindow, sender.GetCongestionWindow()) +} + +func TestCubicSenderMaximumPacketSizeReduction(t *testing.T) { + sender := newTestCubicSender(false) + require.Panics(t, func() { sender.sender.SetMaxDatagramSize(initialMaxDatagramSize - 1) }) +} + +func TestCubicSenderSlowStartsPacketSizeIncrease(t *testing.T) { + var clock mockClock + rttStats := utils.RTTStats{} + const initialMaxCongestionWindow = protocol.MaxCongestionWindowPackets * initialMaxDatagramSize + sender := newCubicSender( + &clock, + &rttStats, + &utils.ConnectionStats{}, + true, + protocol.InitialPacketSize, + initialCongestionWindowPackets*maxDatagramSize, + initialMaxCongestionWindow, + nil, + ) + const packetSize = initialMaxDatagramSize + 100 + sender.SetMaxDatagramSize(packetSize) + for i := 1; i < protocol.MaxCongestionWindowPackets; i++ { + sender.OnPacketAcked(protocol.PacketNumber(i), packetSize, sender.GetCongestionWindow(), clock.Now()) + } + const maxCwnd = protocol.MaxCongestionWindowPackets * packetSize + require.True(t, sender.GetCongestionWindow() > maxCwnd) + require.True(t, sender.GetCongestionWindow() <= maxCwnd+packetSize) +} + +func TestCubicSenderLimitCwndIncreaseInCongestionAvoidance(t *testing.T) { + // Enable Cubic. + var clock mockClock + rttStats := utils.RTTStats{} + sender := newCubicSender( + &clock, + &rttStats, + &utils.ConnectionStats{}, + false, + protocol.InitialPacketSize, + initialCongestionWindowPackets*maxDatagramSize, + MaxCongestionWindow, + nil, + ) + testSender := &testCubicSender{ + sender: sender, + clock: &clock, + rttStats: &rttStats, + } + + numSent := testSender.SendAvailableSendWindow() + + // Make sure we fall out of slow start. + savedCwnd := sender.GetCongestionWindow() + testSender.LoseNPackets(1) + require.Greater(t, savedCwnd, sender.GetCongestionWindow()) + + // Ack the rest of the outstanding packets to get out of recovery. + for i := 1; i < numSent; i++ { + testSender.AckNPackets(1) + } + require.Equal(t, protocol.ByteCount(0), testSender.bytesInFlight) + + savedCwnd = sender.GetCongestionWindow() + testSender.SendAvailableSendWindow() + + // Ack packets until the CWND increases. + for sender.GetCongestionWindow() == savedCwnd { + testSender.AckNPackets(1) + testSender.SendAvailableSendWindow() + } + // Bytes in flight may be larger than the CWND if the CWND isn't an exact + // multiple of the packet sizes being sent. + require.GreaterOrEqual(t, testSender.bytesInFlight, sender.GetCongestionWindow()) + savedCwnd = sender.GetCongestionWindow() + + // Advance time 2 seconds waiting for an ack. + clock.Advance(2 * time.Second) + + // Ack two packets. The CWND should increase by only one packet. + testSender.AckNPackets(2) + require.Equal(t, savedCwnd+maxDatagramSize, sender.GetCongestionWindow()) +} diff --git a/third_party/quic-go/internal/congestion/cubic_test.go b/third_party/quic-go/internal/congestion/cubic_test.go new file mode 100644 index 0000000..174c51a --- /dev/null +++ b/third_party/quic-go/internal/congestion/cubic_test.go @@ -0,0 +1,205 @@ +package congestion + +import ( + "math" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/stretchr/testify/require" +) + +const ( + numConnections uint32 = 2 + nConnectionBeta float32 = (float32(numConnections) - 1 + beta) / float32(numConnections) + nConnectionBetaLastMax float32 = (float32(numConnections) - 1 + betaLastMax) / float32(numConnections) + nConnectionAlpha float32 = 3 * float32(numConnections) * float32(numConnections) * (1 - nConnectionBeta) / (1 + nConnectionBeta) + maxCubicTimeInterval = 30 * time.Millisecond +) + +func renoCwnd(currentCwnd protocol.ByteCount) protocol.ByteCount { + return currentCwnd + protocol.ByteCount(float32(maxDatagramSize)*nConnectionAlpha*float32(maxDatagramSize)/float32(currentCwnd)) +} + +func cubicConvexCwnd(initialCwnd protocol.ByteCount, rtt, elapsedTime time.Duration) protocol.ByteCount { + offset := protocol.ByteCount((elapsedTime+rtt)/time.Microsecond) << 10 / 1000000 + deltaCongestionWindow := 410 * offset * offset * offset * maxDatagramSize >> 40 + return initialCwnd + deltaCongestionWindow +} + +func TestCubicAboveOriginWithTighterBounds(t *testing.T) { + var clock mockClock + cubic := NewCubic(&clock) + cubic.SetNumConnections(int(numConnections)) + + // Convex growth. + const rttMin = 100 * time.Millisecond + const rttMinS = float32(rttMin/time.Millisecond) / 1000.0 + currentCwnd := 10 * maxDatagramSize + initialCwnd := currentCwnd + + clock.Advance(time.Millisecond) + initialTime := clock.Now() + expectedFirstCwnd := renoCwnd(currentCwnd) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, initialTime) + require.Equal(t, expectedFirstCwnd, currentCwnd) + + // Normal TCP phase. + // The maximum number of expected reno RTTs can be calculated by + // finding the point where the cubic curve and the reno curve meet. + maxRenoRtts := int(math.Sqrt(float64(nConnectionAlpha/(0.4*rttMinS*rttMinS*rttMinS))) - 2) + for range maxRenoRtts { + numAcksThisEpoch := int(float32(currentCwnd/maxDatagramSize) / nConnectionAlpha) + + initialCwndThisEpoch := currentCwnd + for range numAcksThisEpoch { + // Call once per ACK. + expectedNextCwnd := renoCwnd(currentCwnd) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + require.Equal(t, expectedNextCwnd, currentCwnd) + } + cwndChangeThisEpoch := currentCwnd - initialCwndThisEpoch + require.InDelta(t, float64(maxDatagramSize), float64(cwndChangeThisEpoch), float64(maxDatagramSize)/2) + clock.Advance(100 * time.Millisecond) + } + + for range 54 { + maxAcksThisEpoch := currentCwnd / maxDatagramSize + interval := time.Duration(100*1000/maxAcksThisEpoch) * time.Microsecond + for range int(maxAcksThisEpoch) { + clock.Advance(interval) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + expectedCwnd := cubicConvexCwnd(initialCwnd, rttMin, clock.Now().Sub(initialTime)) + require.Equal(t, expectedCwnd, currentCwnd) + } + } + expectedCwnd := cubicConvexCwnd(initialCwnd, rttMin, clock.Now().Sub(initialTime)) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + require.Equal(t, expectedCwnd, currentCwnd) +} + +func TestCubicAboveOriginWithFineGrainedCubing(t *testing.T) { + var clock mockClock + cubic := NewCubic(&clock) + cubic.SetNumConnections(int(numConnections)) + + currentCwnd := 1000 * maxDatagramSize + initialCwnd := currentCwnd + rttMin := 100 * time.Millisecond + clock.Advance(time.Millisecond) + initialTime := clock.Now() + + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + clock.Advance(600 * time.Millisecond) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + + for range 100 { + clock.Advance(10 * time.Millisecond) + expectedCwnd := cubicConvexCwnd(initialCwnd, rttMin, clock.Now().Sub(initialTime)) + nextCwnd := cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + require.Equal(t, expectedCwnd, nextCwnd) + require.Greater(t, nextCwnd, currentCwnd) + cwndDelta := nextCwnd - currentCwnd + require.Less(t, cwndDelta, maxDatagramSize/10) + currentCwnd = nextCwnd + } +} + +func TestCubicHandlesPerAckUpdates(t *testing.T) { + var clock mockClock + cubic := NewCubic(&clock) + cubic.SetNumConnections(int(numConnections)) + + initialCwndPackets := 150 + currentCwnd := protocol.ByteCount(initialCwndPackets) * maxDatagramSize + rttMin := 350 * time.Millisecond + + clock.Advance(time.Millisecond) + rCwnd := renoCwnd(currentCwnd) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + initialCwnd := currentCwnd + + maxAcks := int(float32(initialCwndPackets) / nConnectionAlpha) + interval := maxCubicTimeInterval / time.Duration(maxAcks+1) + + clock.Advance(interval) + rCwnd = renoCwnd(rCwnd) + require.Equal(t, currentCwnd, cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now())) + + for range maxAcks - 1 { + clock.Advance(interval) + nextCwnd := cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + rCwnd = renoCwnd(rCwnd) + require.Greater(t, nextCwnd, currentCwnd) + require.Equal(t, rCwnd, nextCwnd) + currentCwnd = nextCwnd + } + + minimumExpectedIncrease := maxDatagramSize * 9 / 10 + require.Greater(t, currentCwnd, initialCwnd+minimumExpectedIncrease) +} + +func TestCubicHandlesLossEvents(t *testing.T) { + var clock mockClock + cubic := NewCubic(&clock) + cubic.SetNumConnections(int(numConnections)) + + rttMin := 100 * time.Millisecond + currentCwnd := 422 * maxDatagramSize + expectedCwnd := renoCwnd(currentCwnd) + + clock.Advance(time.Millisecond) + require.Equal(t, expectedCwnd, cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now())) + + preLossCwnd := currentCwnd + require.Zero(t, cubic.lastMaxCongestionWindow) + expectedCwnd = protocol.ByteCount(float32(currentCwnd) * nConnectionBeta) + require.Equal(t, expectedCwnd, cubic.CongestionWindowAfterPacketLoss(currentCwnd)) + require.Equal(t, preLossCwnd, cubic.lastMaxCongestionWindow) + currentCwnd = expectedCwnd + + preLossCwnd = currentCwnd + expectedCwnd = protocol.ByteCount(float32(currentCwnd) * nConnectionBeta) + require.Equal(t, expectedCwnd, cubic.CongestionWindowAfterPacketLoss(currentCwnd)) + currentCwnd = expectedCwnd + require.Greater(t, preLossCwnd, cubic.lastMaxCongestionWindow) + expectedLastMax := protocol.ByteCount(float32(preLossCwnd) * nConnectionBetaLastMax) + require.Equal(t, expectedLastMax, cubic.lastMaxCongestionWindow) + require.Less(t, expectedCwnd, cubic.lastMaxCongestionWindow) + + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + require.Greater(t, cubic.lastMaxCongestionWindow, currentCwnd) + + currentCwnd = cubic.lastMaxCongestionWindow - 1 + preLossCwnd = currentCwnd + expectedCwnd = protocol.ByteCount(float32(currentCwnd) * nConnectionBeta) + require.Equal(t, expectedCwnd, cubic.CongestionWindowAfterPacketLoss(currentCwnd)) + expectedLastMax = preLossCwnd + require.Equal(t, expectedLastMax, cubic.lastMaxCongestionWindow) +} + +func TestCubicBelowOrigin(t *testing.T) { + var clock mockClock + cubic := NewCubic(&clock) + cubic.SetNumConnections(int(numConnections)) + + rttMin := 100 * time.Millisecond + currentCwnd := 422 * maxDatagramSize + expectedCwnd := renoCwnd(currentCwnd) + + clock.Advance(time.Millisecond) + require.Equal(t, expectedCwnd, cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now())) + + expectedCwnd = protocol.ByteCount(float32(currentCwnd) * nConnectionBeta) + require.Equal(t, expectedCwnd, cubic.CongestionWindowAfterPacketLoss(currentCwnd)) + currentCwnd = expectedCwnd + + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + + for range 40 { + clock.Advance(100 * time.Millisecond) + currentCwnd = cubic.CongestionWindowAfterAck(maxDatagramSize, currentCwnd, rttMin, clock.Now()) + } + expectedCwnd = 553632 * maxDatagramSize / 1460 + require.Equal(t, expectedCwnd, currentCwnd) +} diff --git a/third_party/quic-go/internal/congestion/hybrid_slow_start.go b/third_party/quic-go/internal/congestion/hybrid_slow_start.go new file mode 100644 index 0000000..ac8d3d4 --- /dev/null +++ b/third_party/quic-go/internal/congestion/hybrid_slow_start.go @@ -0,0 +1,112 @@ +package congestion + +import ( + "time" + + "github.com/apernet/quic-go/internal/protocol" +) + +// Note(pwestin): the magic clamping numbers come from the original code in +// tcp_cubic.c. +const hybridStartLowWindow = protocol.ByteCount(16) + +// Number of delay samples for detecting the increase of delay. +const hybridStartMinSamples = uint32(8) + +// Exit slow start if the min rtt has increased by more than 1/8th. +const hybridStartDelayFactorExp = 3 // 2^3 = 8 +// The original paper specifies 2 and 8ms, but those have changed over time. +const ( + hybridStartDelayMinThresholdUs = int64(4000) + hybridStartDelayMaxThresholdUs = int64(16000) +) + +// HybridSlowStart implements the TCP hybrid slow start algorithm +type HybridSlowStart struct { + endPacketNumber protocol.PacketNumber + lastSentPacketNumber protocol.PacketNumber + started bool + currentMinRTT time.Duration + rttSampleCount uint32 + hystartFound bool +} + +// StartReceiveRound is called for the start of each receive round (burst) in the slow start phase. +func (s *HybridSlowStart) StartReceiveRound(lastSent protocol.PacketNumber) { + s.endPacketNumber = lastSent + s.currentMinRTT = 0 + s.rttSampleCount = 0 + s.started = true +} + +// IsEndOfRound returns true if this ack is the last packet number of our current slow start round. +func (s *HybridSlowStart) IsEndOfRound(ack protocol.PacketNumber) bool { + return s.endPacketNumber < ack +} + +// ShouldExitSlowStart should be called on every new ack frame, since a new +// RTT measurement can be made then. +// rtt: the RTT for this ack packet. +// minRTT: is the lowest delay (RTT) we have seen during the session. +// congestionWindow: the congestion window in packets. +func (s *HybridSlowStart) ShouldExitSlowStart(latestRTT time.Duration, minRTT time.Duration, congestionWindow protocol.ByteCount) bool { + if !s.started { + // Time to start the hybrid slow start. + s.StartReceiveRound(s.lastSentPacketNumber) + } + if s.hystartFound { + return true + } + // Second detection parameter - delay increase detection. + // Compare the minimum delay (s.currentMinRTT) of the current + // burst of packets relative to the minimum delay during the session. + // Note: we only look at the first few(8) packets in each burst, since we + // only want to compare the lowest RTT of the burst relative to previous + // bursts. + s.rttSampleCount++ + if s.rttSampleCount <= hybridStartMinSamples { + if s.currentMinRTT == 0 || s.currentMinRTT > latestRTT { + s.currentMinRTT = latestRTT + } + } + // We only need to check this once per round. + if s.rttSampleCount == hybridStartMinSamples { + // Divide minRTT by 8 to get a rtt increase threshold for exiting. + minRTTincreaseThresholdUs := int64(minRTT / time.Microsecond >> hybridStartDelayFactorExp) + // Ensure the rtt threshold is never less than 2ms or more than 16ms. + minRTTincreaseThresholdUs = min(minRTTincreaseThresholdUs, hybridStartDelayMaxThresholdUs) + minRTTincreaseThreshold := time.Duration(max(minRTTincreaseThresholdUs, hybridStartDelayMinThresholdUs)) * time.Microsecond + + if s.currentMinRTT > (minRTT + minRTTincreaseThreshold) { + s.hystartFound = true + } + } + // Exit from slow start if the cwnd is greater than 16 and + // increasing delay is found. + return congestionWindow >= hybridStartLowWindow && s.hystartFound +} + +// OnPacketSent is called when a packet was sent +func (s *HybridSlowStart) OnPacketSent(packetNumber protocol.PacketNumber) { + s.lastSentPacketNumber = packetNumber +} + +// OnPacketAcked gets invoked after ShouldExitSlowStart, so it's best to end +// the round when the final packet of the burst is received and start it on +// the next incoming ack. +func (s *HybridSlowStart) OnPacketAcked(ackedPacketNumber protocol.PacketNumber) { + if s.IsEndOfRound(ackedPacketNumber) { + s.started = false + } +} + +// Started returns true if started +func (s *HybridSlowStart) Started() bool { + return s.started +} + +// Restart the slow start phase +func (s *HybridSlowStart) Restart() { + s.started = false + s.hystartFound = false +} diff --git a/third_party/quic-go/internal/congestion/hybrid_slow_start_test.go b/third_party/quic-go/internal/congestion/hybrid_slow_start_test.go new file mode 100644 index 0000000..6003de8 --- /dev/null +++ b/third_party/quic-go/internal/congestion/hybrid_slow_start_test.go @@ -0,0 +1,68 @@ +package congestion + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestHybridSlowStartSimpleCase(t *testing.T) { + slowStart := HybridSlowStart{} + + packetNumber := protocol.PacketNumber(1) + endPacketNumber := protocol.PacketNumber(3) + slowStart.StartReceiveRound(endPacketNumber) + + packetNumber++ + require.False(t, slowStart.IsEndOfRound(packetNumber)) + + // Test duplicates. + require.False(t, slowStart.IsEndOfRound(packetNumber)) + + packetNumber++ + require.False(t, slowStart.IsEndOfRound(packetNumber)) + packetNumber++ + require.True(t, slowStart.IsEndOfRound(packetNumber)) + + // Test without a new registered end_packet_number; + packetNumber++ + require.True(t, slowStart.IsEndOfRound(packetNumber)) + + endPacketNumber = 20 + slowStart.StartReceiveRound(endPacketNumber) + for packetNumber < endPacketNumber { + packetNumber++ + require.False(t, slowStart.IsEndOfRound(packetNumber)) + } + packetNumber++ + require.True(t, slowStart.IsEndOfRound(packetNumber)) +} + +func TestHybridSlowStartWithDelay(t *testing.T) { + slowStart := HybridSlowStart{} + const rtt = 60 * time.Millisecond + // We expect to detect the increase at +1/8 of the RTT; hence at a typical + // RTT of 60ms the detection will happen at 67.5 ms. + const hybridStartMinSamples = 8 // Number of acks required to trigger. + + endPacketNumber := protocol.PacketNumber(1) + endPacketNumber++ + slowStart.StartReceiveRound(endPacketNumber) + + // Will not trigger since our lowest RTT in our burst is the same as the long + // term RTT provided. + for n := range hybridStartMinSamples { + require.False(t, slowStart.ShouldExitSlowStart(rtt+time.Duration(n)*time.Millisecond, rtt, 100)) + } + endPacketNumber++ + slowStart.StartReceiveRound(endPacketNumber) + for n := 1; n < hybridStartMinSamples; n++ { + require.False(t, slowStart.ShouldExitSlowStart(rtt+(time.Duration(n)+10)*time.Millisecond, rtt, 100)) + } + // Expect to trigger since all packets in this burst was above the long term + // RTT provided. + require.True(t, slowStart.ShouldExitSlowStart(rtt+10*time.Millisecond, rtt, 100)) +} diff --git a/third_party/quic-go/internal/congestion/interface.go b/third_party/quic-go/internal/congestion/interface.go new file mode 100644 index 0000000..6f56448 --- /dev/null +++ b/third_party/quic-go/internal/congestion/interface.go @@ -0,0 +1,33 @@ +package congestion + +import ( + "github.com/apernet/quic-go/congestion" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +// A SendAlgorithm performs congestion control +type SendAlgorithm interface { + TimeUntilSend(bytesInFlight protocol.ByteCount) monotime.Time + HasPacingBudget(now monotime.Time) bool + OnPacketSent(sentTime monotime.Time, bytesInFlight protocol.ByteCount, packetNumber protocol.PacketNumber, bytes protocol.ByteCount, isRetransmittable bool) + CanSend(bytesInFlight protocol.ByteCount) bool + MaybeExitSlowStart() + OnPacketAcked(number protocol.PacketNumber, ackedBytes protocol.ByteCount, priorInFlight protocol.ByteCount, eventTime monotime.Time) + OnCongestionEvent(number protocol.PacketNumber, lostBytes protocol.ByteCount, priorInFlight protocol.ByteCount) + OnRetransmissionTimeout(packetsRetransmitted bool) + SetMaxDatagramSize(protocol.ByteCount) +} + +type SendAlgorithmEx interface { + SendAlgorithm + OnCongestionEventEx(priorInFlight protocol.ByteCount, eventTime monotime.Time, ackedPackets []congestion.AckedPacketInfo, lostPackets []congestion.LostPacketInfo) +} + +// A SendAlgorithmWithDebugInfos is a SendAlgorithm that exposes some debug infos +type SendAlgorithmWithDebugInfos interface { + SendAlgorithm + InSlowStart() bool + InRecovery() bool + GetCongestionWindow() protocol.ByteCount +} diff --git a/third_party/quic-go/internal/congestion/pacer.go b/third_party/quic-go/internal/congestion/pacer.go new file mode 100644 index 0000000..3862111 --- /dev/null +++ b/third_party/quic-go/internal/congestion/pacer.go @@ -0,0 +1,110 @@ +package congestion + +import ( + "math" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" +) + +const maxBurstSizePackets = 10 + +// The pacer implements a token bucket pacing algorithm. +type pacer struct { + budgetAtLastSent protocol.ByteCount + maxDatagramSize protocol.ByteCount + lastSentTime monotime.Time + adjustedBandwidth func() uint64 // in bytes/s +} + +func newPacer(getBandwidth func() Bandwidth) *pacer { + p := &pacer{ + maxDatagramSize: initialMaxDatagramSize, + adjustedBandwidth: func() uint64 { + // Bandwidth is in bits/s. We need the value in bytes/s. + bw := uint64(getBandwidth() / BytesPerSecond) + // Use a slightly higher value than the actual measured bandwidth. + // RTT variations then won't result in under-utilization of the congestion window. + // Ultimately, this will result in sending packets as acknowledgments are received rather than when timers fire, + // provided the congestion window is fully utilized and acknowledgments arrive at regular intervals. + return bw * 5 / 4 + }, + } + p.budgetAtLastSent = p.maxBurstSize() + return p +} + +func (p *pacer) SentPacket(sendTime monotime.Time, size protocol.ByteCount) { + budget := p.Budget(sendTime) + if size >= budget { + p.budgetAtLastSent = 0 + } else { + p.budgetAtLastSent = budget - size + } + p.lastSentTime = sendTime +} + +func (p *pacer) Budget(now monotime.Time) protocol.ByteCount { + if p.lastSentTime.IsZero() { + return p.maxBurstSize() + } + delta := now.Sub(p.lastSentTime) + var added protocol.ByteCount + if delta > 0 { + added = p.timeScaledBandwidth(uint64(delta.Nanoseconds())) + } + budget := p.budgetAtLastSent + added + if added > 0 && budget < p.budgetAtLastSent { + budget = protocol.MaxByteCount + } + return min(p.maxBurstSize(), budget) +} + +func (p *pacer) maxBurstSize() protocol.ByteCount { + return max( + p.timeScaledBandwidth(uint64((protocol.MinPacingDelay + protocol.TimerGranularity).Nanoseconds())), + maxBurstSizePackets*p.maxDatagramSize, + ) +} + +// timeScaledBandwidth calculates the number of bytes that may be sent within +// a given time interval (ns nanoseconds), based on the current bandwidth estimate. +// It caps the scaled value to the maximum allowed burst and handles overflows. +func (p *pacer) timeScaledBandwidth(ns uint64) protocol.ByteCount { + bw := p.adjustedBandwidth() + if bw == 0 { + return 0 + } + const nsPerSecond = 1e9 + maxBurst := maxBurstSizePackets * p.maxDatagramSize + var scaled protocol.ByteCount + if ns > math.MaxUint64/bw { + scaled = maxBurst + } else { + scaled = protocol.ByteCount(bw * ns / nsPerSecond) + } + return scaled +} + +// TimeUntilSend returns when the next packet should be sent. +// It returns zero if a packet can be sent immediately. +func (p *pacer) TimeUntilSend() monotime.Time { + if p.budgetAtLastSent >= p.maxDatagramSize { + return 0 + } + diff := 1e9 * uint64(p.maxDatagramSize-p.budgetAtLastSent) + bw := p.adjustedBandwidth() + // We might need to round up this value. + // Otherwise, we might have a budget (slightly) smaller than the datagram size when the timer expires. + d := diff / bw + // this is effectively a math.Ceil, but using only integer math + if diff%bw > 0 { + d++ + } + return p.lastSentTime.Add(max(protocol.MinPacingDelay, time.Duration(d)*time.Nanosecond)) +} + +func (p *pacer) SetMaxDatagramSize(s protocol.ByteCount) { + p.maxDatagramSize = s +} diff --git a/third_party/quic-go/internal/congestion/pacer_test.go b/third_party/quic-go/internal/congestion/pacer_test.go new file mode 100644 index 0000000..0ed0a25 --- /dev/null +++ b/third_party/quic-go/internal/congestion/pacer_test.go @@ -0,0 +1,154 @@ +package congestion + +import ( + "math" + "math/rand/v2" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestPacerPacing(t *testing.T) { + bandwidth := 50 * initialMaxDatagramSize // 50 full-size packets per second + p := newPacer(func() Bandwidth { return Bandwidth(bandwidth) * BytesPerSecond * 4 / 5 }) + now := monotime.Now() + require.Zero(t, p.TimeUntilSend()) + budget := p.Budget(now) + require.Equal(t, maxBurstSizePackets*initialMaxDatagramSize, budget) + + // consume the initial budget by sending packets + for budget > 0 { + require.Zero(t, p.TimeUntilSend()) + require.Equal(t, budget, p.Budget(now)) + p.SentPacket(now, initialMaxDatagramSize) + budget -= initialMaxDatagramSize + } + + // now packets are being paced + for range 5 { + require.Zero(t, p.Budget(now)) + nextPacket := p.TimeUntilSend() + require.NotZero(t, nextPacket) + require.Equal(t, time.Second/50, nextPacket.Sub(now)) + now = nextPacket + p.SentPacket(now, initialMaxDatagramSize) + } + + nextPacket := p.TimeUntilSend() + require.Equal(t, time.Second/50, nextPacket.Sub(now)) + // send this packet a bit later, simulating timer delay + p.SentPacket(nextPacket.Add(time.Millisecond), initialMaxDatagramSize) + // the next packet should be paced again, without a delay + require.Equal(t, time.Second/50, p.TimeUntilSend().Sub(nextPacket)) + + // now send a half-size packet + now = p.TimeUntilSend() + p.SentPacket(now, initialMaxDatagramSize/2) + require.Equal(t, initialMaxDatagramSize/2, p.Budget(now)) + require.Equal(t, time.Second/100, p.TimeUntilSend().Sub(now)) + p.SentPacket(p.TimeUntilSend(), initialMaxDatagramSize/2) + + now = p.TimeUntilSend() + // budget accumulates if no packets are sent for a while + // we should have accumulated budget to send a burst now + require.Equal(t, 5*initialMaxDatagramSize, p.Budget(now.Add(4*time.Second/50))) + // but the budget is capped at the max burst size + require.Equal(t, maxBurstSizePackets*initialMaxDatagramSize, p.Budget(now.Add(time.Hour))) + p.SentPacket(now, initialMaxDatagramSize) + require.Zero(t, p.Budget(now)) + + // reduce the bandwidth + bandwidth = 10 * initialMaxDatagramSize // 10 full-size packets per second + require.Equal(t, time.Second/10, p.TimeUntilSend().Sub(now)) +} + +func TestPacerUpdatePacketSize(t *testing.T) { + const bandwidth = 50 * initialMaxDatagramSize // 50 full-size packets per second + p := newPacer(func() Bandwidth { return Bandwidth(bandwidth) * BytesPerSecond * 4 / 5 }) + + // consume the initial budget by sending packets + now := monotime.Now() + for p.Budget(now) > 0 { + p.SentPacket(now, initialMaxDatagramSize) + } + + require.Equal(t, time.Second/50, p.TimeUntilSend().Sub(now)) + // Double the packet size. We now need to wait twice as long to send the next packet. + const newDatagramSize = 2 * initialMaxDatagramSize + p.SetMaxDatagramSize(newDatagramSize) + require.Equal(t, 2*time.Second/50, p.TimeUntilSend().Sub(now)) + + // check that the maximum burst size is updated + require.Equal(t, maxBurstSizePackets*newDatagramSize, p.Budget(now.Add(time.Hour))) +} + +func TestPacerFastPacing(t *testing.T) { + const bandwidth = 10000 * initialMaxDatagramSize // 10,000 full-size packets per second + p := newPacer(func() Bandwidth { return Bandwidth(bandwidth) * BytesPerSecond * 4 / 5 }) + + // consume the initial budget by sending packets + now := monotime.Now() + for p.Budget(now) > 0 { + p.SentPacket(now, initialMaxDatagramSize) + } + + // If we were pacing by packet, we'd expect the next packet to send in 1/10ms. + // However, we don't want to arm the pacing timer for less than 1ms, + // so we wait for 1ms, and then send 10 packets in a burst. + require.Equal(t, time.Millisecond, p.TimeUntilSend().Sub(now)) + require.Equal(t, 10*initialMaxDatagramSize, p.Budget(now.Add(time.Millisecond))) + + now = now.Add(time.Millisecond) + for range 10 { + require.NotZero(t, p.Budget(now)) + p.SentPacket(now, initialMaxDatagramSize) + } + require.Zero(t, p.Budget(now)) + require.Equal(t, time.Millisecond, p.TimeUntilSend().Sub(now)) +} + +func TestPacerNoOverflows(t *testing.T) { + p := newPacer(func() Bandwidth { return math.MaxUint64 }) + now := monotime.Now() + p.SentPacket(now, initialMaxDatagramSize) + for range 100000 { + require.NotZero(t, p.Budget(now.Add(time.Duration(rand.Int64N(math.MaxInt64))))) + } + + burstCount := 1 + for p.Budget(now) > 0 { + burstCount++ + p.SentPacket(now, initialMaxDatagramSize) + } + require.Equal(t, maxBurstSizePackets, burstCount) + require.Zero(t, p.Budget(now)) + + next := p.TimeUntilSend() + require.Equal(t, next.Sub(now), protocol.MinPacingDelay) + require.Greater(t, p.Budget(next), initialMaxDatagramSize) +} + +func BenchmarkPacer(b *testing.B) { + const bandwidth = 50 * initialMaxDatagramSize // 50 full-size packets per second + p := newPacer(func() Bandwidth { return Bandwidth(bandwidth) * BytesPerSecond * 4 / 5 }) + + now := monotime.Now() + + var i int + for b.Loop() { + i++ + for p.Budget(now) > 0 { + p.SentPacket(now, initialMaxDatagramSize) + } + next := p.TimeUntilSend() + if i%2 == 0 { + now = next + } else { + now = now.Add(100 * time.Millisecond) + } + } +} diff --git a/third_party/quic-go/internal/handshake/aead.go b/third_party/quic-go/internal/handshake/aead.go new file mode 100644 index 0000000..5523e70 --- /dev/null +++ b/third_party/quic-go/internal/handshake/aead.go @@ -0,0 +1,91 @@ +package handshake + +import ( + "crypto/cipher" + "encoding/binary" + + "github.com/apernet/quic-go/internal/protocol" +) + +func createAEAD(suite cipherSuite, trafficSecret []byte, v protocol.Version) cipher.AEAD { + keyLabel := hkdfLabelKeyV1 + ivLabel := hkdfLabelIVV1 + if v == protocol.Version2 { + keyLabel = hkdfLabelKeyV2 + ivLabel = hkdfLabelIVV2 + } + key := hkdfExpandLabel(suite.Hash, trafficSecret, []byte{}, keyLabel, suite.KeyLen) + iv := hkdfExpandLabel(suite.Hash, trafficSecret, []byte{}, ivLabel, suite.IVLen()) + return suite.AEAD(key, iv) +} + +type longHeaderSealer struct { + aead cipher.AEAD + headerProtector headerProtector + nonceBuf [8]byte +} + +var _ LongHeaderSealer = &longHeaderSealer{} + +func newLongHeaderSealer(aead cipher.AEAD, headerProtector headerProtector) LongHeaderSealer { + if aead.NonceSize() != 8 { + panic("unexpected nonce size") + } + return &longHeaderSealer{ + aead: aead, + headerProtector: headerProtector, + } +} + +func (s *longHeaderSealer) Seal(dst, src []byte, pn protocol.PacketNumber, ad []byte) []byte { + binary.BigEndian.PutUint64(s.nonceBuf[:], uint64(pn)) + return s.aead.Seal(dst, s.nonceBuf[:], src, ad) +} + +func (s *longHeaderSealer) EncryptHeader(sample []byte, firstByte *byte, pnBytes []byte) { + s.headerProtector.EncryptHeader(sample, firstByte, pnBytes) +} + +func (s *longHeaderSealer) Overhead() int { + return s.aead.Overhead() +} + +type longHeaderOpener struct { + aead cipher.AEAD + headerProtector headerProtector + highestRcvdPN protocol.PacketNumber // highest packet number received (which could be successfully unprotected) + + // use a single array to avoid allocations + nonceBuf [8]byte +} + +var _ LongHeaderOpener = &longHeaderOpener{} + +func newLongHeaderOpener(aead cipher.AEAD, headerProtector headerProtector) LongHeaderOpener { + if aead.NonceSize() != 8 { + panic("unexpected nonce size") + } + return &longHeaderOpener{ + aead: aead, + headerProtector: headerProtector, + } +} + +func (o *longHeaderOpener) DecodePacketNumber(wirePN protocol.PacketNumber, wirePNLen protocol.PacketNumberLen) protocol.PacketNumber { + return protocol.DecodePacketNumber(wirePNLen, o.highestRcvdPN, wirePN) +} + +func (o *longHeaderOpener) Open(dst, src []byte, pn protocol.PacketNumber, ad []byte) ([]byte, error) { + binary.BigEndian.PutUint64(o.nonceBuf[:], uint64(pn)) + dec, err := o.aead.Open(dst, o.nonceBuf[:], src, ad) + if err == nil { + o.highestRcvdPN = max(o.highestRcvdPN, pn) + } else { + err = ErrDecryptionFailed + } + return dec, err +} + +func (o *longHeaderOpener) DecryptHeader(sample []byte, firstByte *byte, pnBytes []byte) { + o.headerProtector.DecryptHeader(sample, firstByte, pnBytes) +} diff --git a/third_party/quic-go/internal/handshake/aead_test.go b/third_party/quic-go/internal/handshake/aead_test.go new file mode 100644 index 0000000..f0cb0ac --- /dev/null +++ b/third_party/quic-go/internal/handshake/aead_test.go @@ -0,0 +1,108 @@ +package handshake + +import ( + "crypto/rand" + "crypto/tls" + "fmt" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func getSealerAndOpener(t *testing.T, cs cipherSuite, v protocol.Version) (LongHeaderSealer, LongHeaderOpener) { + t.Helper() + trafficSecret := make([]byte, cs.Hash.Size()) + hpKey := make([]byte, 16) + rand.Read(trafficSecret) + rand.Read(hpKey) + aead := createAEAD(cs, trafficSecret, v) + return newLongHeaderSealer(aead, newHeaderProtector(cs, hpKey, true, v)), + newLongHeaderOpener(aead, newHeaderProtector(cs, hpKey, true, v)) +} + +func TestEncryptAndDecryptMessage(t *testing.T) { + for _, v := range []protocol.Version{protocol.Version1, protocol.Version2} { + for _, cs := range cipherSuites { + t.Run(fmt.Sprintf("QUIC %s/%s", v, tls.CipherSuiteName(cs.ID)), func(t *testing.T) { + sealer, opener := getSealerAndOpener(t, cs, v) + msg := []byte("Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.") + ad := []byte("Donec in velit neque.") + + encrypted := sealer.Seal(nil, msg, 0x1337, ad) + + opened, err := opener.Open(nil, encrypted, 0x1337, ad) + require.NoError(t, err) + require.Equal(t, msg, opened) + + // incorrect associated data + _, err = opener.Open(nil, encrypted, 0x1337, []byte("wrong ad")) + require.Equal(t, ErrDecryptionFailed, err) + + // incorrect packet number + _, err = opener.Open(nil, encrypted, 0x42, ad) + require.Equal(t, ErrDecryptionFailed, err) + }) + } + } +} + +func TestDecodePacketNumber(t *testing.T) { + msg := []byte("Lorem ipsum dolor sit amet") + ad := []byte("Donec in velit neque.") + + sealer, opener := getSealerAndOpener(t, getCipherSuite(tls.TLS_AES_128_GCM_SHA256), protocol.Version1) + encrypted := sealer.Seal(nil, msg, 0x1337, ad) + + // can't decode the packet number if encryption failed + _, err := opener.Open(nil, encrypted[:len(encrypted)-1], 0x1337, ad) + require.Error(t, err) + require.Equal(t, protocol.PacketNumber(0x38), opener.DecodePacketNumber(0x38, protocol.PacketNumberLen1)) + + _, err = opener.Open(nil, encrypted, 0x1337, ad) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumber(0x1338), opener.DecodePacketNumber(0x38, protocol.PacketNumberLen1)) +} + +func TestEncryptAndDecryptHeader(t *testing.T) { + for _, v := range []protocol.Version{protocol.Version1, protocol.Version2} { + t.Run("QUIC "+v.String(), func(t *testing.T) { + for _, cs := range cipherSuites { + t.Run(tls.CipherSuiteName(cs.ID), func(t *testing.T) { + testEncryptAndDecryptHeader(t, cs, v) + }) + } + }) + } +} + +func testEncryptAndDecryptHeader(t *testing.T, cs cipherSuite, v protocol.Version) { + sealer, opener := getSealerAndOpener(t, cs, v) + var lastFourBitsDifferent int + + for range 100 { + sample := make([]byte, 16) + rand.Read(sample) + header := []byte{0xb5, 1, 2, 3, 4, 5, 6, 7, 8, 0xde, 0xad, 0xbe, 0xef} + sealer.EncryptHeader(sample, &header[0], header[9:13]) + if header[0]&0xf != 0xb5&0xf { + lastFourBitsDifferent++ + } + require.Equal(t, byte(0xb5&0xf0), header[0]&0xf0) + require.Equal(t, []byte{1, 2, 3, 4, 5, 6, 7, 8}, header[1:9]) + require.NotEqual(t, []byte{0xde, 0xad, 0xbe, 0xef}, header[9:13]) + opener.DecryptHeader(sample, &header[0], header[9:13]) + require.Equal(t, []byte{0xb5, 1, 2, 3, 4, 5, 6, 7, 8, 0xde, 0xad, 0xbe, 0xef}, header) + } + require.Greater(t, lastFourBitsDifferent, 75) + + // decryption failure with different sample + header := []byte{0xb5, 1, 2, 3, 4, 5, 6, 7, 8, 0xde, 0xad, 0xbe, 0xef} + sample := make([]byte, 16) + rand.Read(sample) + sealer.EncryptHeader(sample, &header[0], header[9:13]) + rand.Read(sample) // use a different sample + opener.DecryptHeader(sample, &header[0], header[9:13]) + require.NotEqual(t, []byte{0xb5, 1, 2, 3, 4, 5, 6, 7, 8, 0xde, 0xad, 0xbe, 0xef}, header) +} diff --git a/third_party/quic-go/internal/handshake/chrome_client_hello.go b/third_party/quic-go/internal/handshake/chrome_client_hello.go new file mode 100644 index 0000000..3eaf739 --- /dev/null +++ b/third_party/quic-go/internal/handshake/chrome_client_hello.go @@ -0,0 +1,76 @@ +package handshake + +import ( + utls "github.com/refraction-networking/utls" +) + +// chromeQUICClientHelloSpec returns the TLS ClientHello Chrome sends over QUIC. +// +// This is deliberately NOT one of uTLS's HelloChrome_* presets. Those are +// TLS-over-TCP (h2) fingerprints, and the QUIC ClientHello is a materially +// different message; using a TCP preset over QUIC produces a combination +// matching no real client, which is worse than not trying at all. +// +// The extension order is permuted per connection, hence the shuffle. +// +// alpn comes from the caller's tls.Config rather than being hardcoded: +// advertising a protocol the peer doesn't speak breaks the connection outright, +// which is worse than an imperfect fingerprint. The match is exact only for the +// protocol Chrome would negotiate. +func chromeQUICClientHelloSpec(alpn []string) *utls.ClientHelloSpec { + if len(alpn) == 0 { + alpn = []string{"h3"} + } + return &utls.ClientHelloSpec{ + // No GREASE suite here, unlike the TCP hello. + CipherSuites: []uint16{ + utls.TLS_AES_128_GCM_SHA256, + utls.TLS_AES_256_GCM_SHA384, + utls.TLS_CHACHA20_POLY1305_SHA256, + }, + CompressionMethods: []byte{0x00}, + Extensions: utls.ShuffleChromeTLSExtensions([]utls.TLSExtension{ + &utls.SNIExtension{}, + // No GREASE curve, unlike the TCP hello. + &utls.SupportedCurvesExtension{Curves: []utls.CurveID{ + utls.X25519MLKEM768, + utls.X25519, + utls.CurveP256, + utls.CurveP384, + }}, + &utls.SignatureAlgorithmsExtension{SupportedSignatureAlgorithms: []utls.SignatureScheme{ + utls.ECDSAWithP256AndSHA256, + utls.PSSWithSHA256, + utls.PKCS1WithSHA256, + utls.ECDSAWithP384AndSHA384, + utls.PSSWithSHA384, + utls.PKCS1WithSHA384, + utls.PSSWithSHA512, + utls.PKCS1WithSHA512, + utls.PKCS1WithSHA1, + }}, + &utls.ALPNExtension{AlpnProtocols: alpn}, + &utls.UtlsCompressCertExtension{Algorithms: []utls.CertCompressionAlgo{ + utls.CertCompressionBrotli, + }}, + // TLS 1.3 only, again with no GREASE version. + &utls.SupportedVersionsExtension{Versions: []uint16{utls.VersionTLS13}}, + &utls.PSKKeyExchangeModesExtension{Modes: []uint8{utls.PskModeDHE}}, + // The ML-KEM share is what pushes the ClientHello past a single packet, + // giving the characteristic two Initial datagrams. + &utls.KeyShareExtension{KeyShares: []utls.KeyShare{ + {Group: utls.X25519MLKEM768}, + {Group: utls.X25519}, + }}, + // Populated from quic-go's own marshalled parameters; see + // utlsQUICConn.SetTransportParameters. + &utls.QUICTransportParametersExtension{}, + // ALPS at the newer codepoint; uTLS's presets use the older one. + &utls.ApplicationSettingsExtensionNew{SupportedProtocols: alpn}, + // An ECH extension is always present: real when a config is available + // from DNS, GREASE otherwise. GREASE is right here, and is structurally + // indistinguishable from the real thing without decrypting it. + utls.BoringGREASEECH(), + }), + } +} diff --git a/third_party/quic-go/internal/handshake/cipher_suite.go b/third_party/quic-go/internal/handshake/cipher_suite.go new file mode 100644 index 0000000..98112a8 --- /dev/null +++ b/third_party/quic-go/internal/handshake/cipher_suite.go @@ -0,0 +1,114 @@ +package handshake + +import ( + "crypto" + "crypto/aes" + "crypto/cipher" + "crypto/fips140" + "crypto/tls" + "fmt" + + "golang.org/x/crypto/chacha20poly1305" +) + +// These cipher suite implementations are copied from the standard library crypto/tls package. + +const aeadNonceLength = 12 + +type cipherSuite struct { + ID uint16 + Hash crypto.Hash + KeyLen int + AEAD func(key, nonceMask []byte) cipher.AEAD +} + +func (s cipherSuite) IVLen() int { return aeadNonceLength } + +func getCipherSuite(id uint16) cipherSuite { + switch id { + case tls.TLS_AES_128_GCM_SHA256: + return cipherSuite{ID: tls.TLS_AES_128_GCM_SHA256, Hash: crypto.SHA256, KeyLen: 16, AEAD: aeadAESGCMTLS13} + case tls.TLS_CHACHA20_POLY1305_SHA256: + // The usual convention is to only panic on fips140.Enforced (and not on fips140.Enabled), + // but this function panics in the default case anyway, so we might as well panic here. + if fips140.Enabled() { + panic("tls: TLS_CHACHA20_POLY1305_SHA256 is not allowed in FIPS 140-3 mode") + } + return cipherSuite{ID: tls.TLS_CHACHA20_POLY1305_SHA256, Hash: crypto.SHA256, KeyLen: 32, AEAD: aeadChaCha20Poly1305} + case tls.TLS_AES_256_GCM_SHA384: + return cipherSuite{ID: tls.TLS_AES_256_GCM_SHA384, Hash: crypto.SHA384, KeyLen: 32, AEAD: aeadAESGCMTLS13} + default: + panic(fmt.Sprintf("unknown cypher suite: %d", id)) + } +} + +func aeadAESGCMTLS13(key, nonceMask []byte) cipher.AEAD { + if fips140.Enabled() { + return aeadAESGCMTLS13FIPS140(key, nonceMask) + } + + if len(nonceMask) != aeadNonceLength { + panic("tls: internal error: wrong nonce length") + } + aes, err := aes.NewCipher(key) + if err != nil { + panic(err) + } + aead, err := cipher.NewGCM(aes) + if err != nil { + panic(err) + } + + ret := &xorNonceAEAD{aead: aead} + copy(ret.nonceMask[:], nonceMask) + return ret +} + +func aeadChaCha20Poly1305(key, nonceMask []byte) cipher.AEAD { + if len(nonceMask) != aeadNonceLength { + panic("tls: internal error: wrong nonce length") + } + aead, err := chacha20poly1305.New(key) + if err != nil { + panic(err) + } + + ret := &xorNonceAEAD{aead: aead} + copy(ret.nonceMask[:], nonceMask) + return ret +} + +// xorNonceAEAD wraps an AEAD by XORing in a fixed pattern to the nonce +// before each call. +type xorNonceAEAD struct { + nonceMask [aeadNonceLength]byte + aead cipher.AEAD +} + +func (f *xorNonceAEAD) NonceSize() int { return 8 } // 64-bit sequence number +func (f *xorNonceAEAD) Overhead() int { return f.aead.Overhead() } +func (f *xorNonceAEAD) explicitNonceLen() int { return 0 } + +func (f *xorNonceAEAD) Seal(out, nonce, plaintext, additionalData []byte) []byte { + for i, b := range nonce { + f.nonceMask[4+i] ^= b + } + result := f.aead.Seal(out, f.nonceMask[:], plaintext, additionalData) + for i, b := range nonce { + f.nonceMask[4+i] ^= b + } + + return result +} + +func (f *xorNonceAEAD) Open(out, nonce, ciphertext, additionalData []byte) ([]byte, error) { + for i, b := range nonce { + f.nonceMask[4+i] ^= b + } + result, err := f.aead.Open(out, f.nonceMask[:], ciphertext, additionalData) + for i, b := range nonce { + f.nonceMask[4+i] ^= b + } + + return result, err +} diff --git a/third_party/quic-go/internal/handshake/cipher_suite_fips140.go b/third_party/quic-go/internal/handshake/cipher_suite_fips140.go new file mode 100644 index 0000000..2248cba --- /dev/null +++ b/third_party/quic-go/internal/handshake/cipher_suite_fips140.go @@ -0,0 +1,48 @@ +package handshake + +import ( + "crypto/cipher" + _ "unsafe" // for go:linkname +) + +// Reaching into crypto/tls is a bit of a hack, but it's the only way to get the FIPS 140 +// compliant AEAD, because the standard library doesn't yet expose the NewGCMForQUIC constructor +// added in https://go-review.googlesource.com/c/go/+/723760. +// See https://github.com/golang/go/issues/79219 for details. +// +// Once the standard library exposes the necessary constructors, we can use a shared code path +// for both FIPS 140 and non-FIPS 140 modes. +// +//go:linkname cryptoTLSAEAD_AESGCMTLS13 crypto/tls.aeadAESGCMTLS13 +func cryptoTLSAEAD_AESGCMTLS13(key, nonceMask []byte) cipher.AEAD + +func aeadAESGCMTLS13FIPS140(key, nonceMask []byte) cipher.AEAD { + return &tls13AESGCMAEADFIPS140{aead: cryptoTLSAEAD_AESGCMTLS13(key, nonceMask)} +} + +type tls13AESGCMAEADFIPS140 struct { + aead cipher.AEAD + primedSeal bool +} + +func (f *tls13AESGCMAEADFIPS140) NonceSize() int { return f.aead.NonceSize() } +func (f *tls13AESGCMAEADFIPS140) Overhead() int { return f.aead.Overhead() } + +func (f *tls13AESGCMAEADFIPS140) Seal(out, nonce, plaintext, additionalData []byte) []byte { + if !f.primedSeal { + f.primedSeal = true + if nonce[0]|nonce[1]|nonce[2]|nonce[3]|nonce[4]|nonce[5]|nonce[6]|nonce[7] != 0 { + // Go's TLS 1.3 AES-GCM AEAD learns the XOR mask from the first Seal + // call and enforces monotonically increasing packet numbers after that. + // QUIC packet numbers don't reset on key updates, so prime it with + // packet number 0 before the first real, non-zero packet number. + var zeroNonce [8]byte + f.aead.Seal(nil, zeroNonce[:], nil, nil) + } + } + return f.aead.Seal(out, nonce, plaintext, additionalData) +} + +func (f *tls13AESGCMAEADFIPS140) Open(out, nonce, ciphertext, additionalData []byte) ([]byte, error) { + return f.aead.Open(out, nonce, ciphertext, additionalData) +} diff --git a/third_party/quic-go/internal/handshake/crypto_setup.go b/third_party/quic-go/internal/handshake/crypto_setup.go new file mode 100644 index 0000000..f0a3607 --- /dev/null +++ b/third_party/quic-go/internal/handshake/crypto_setup.go @@ -0,0 +1,730 @@ +package handshake + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "strings" + "sync/atomic" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/quicvarint" +) + +type quicVersionContextKey struct{} + +var QUICVersionContextKey = &quicVersionContextKey{} + +const clientSessionStateRevision = 5 + +type cryptoSetup struct { + tlsConf *tls.Config + conn tlsQUICConn + + events []Event + + version protocol.Version + + ourParams *wire.TransportParameters + peerParams *wire.TransportParameters + + zeroRTTParameters *wire.TransportParameters + allow0RTT bool + + rttStats *utils.RTTStats + + qlogger qlogwriter.Recorder + logger utils.Logger + + perspective protocol.Perspective + + handshakeCompleteTime time.Time + + zeroRTTOpener LongHeaderOpener // only set for the server + zeroRTTSealer LongHeaderSealer // only set for the client + + initialOpener LongHeaderOpener + initialSealer LongHeaderSealer + + handshakeOpener LongHeaderOpener + handshakeSealer LongHeaderSealer + + used0RTT atomic.Bool + + aead *updatableAEAD + has1RTTSealer bool + has1RTTOpener bool +} + +var _ CryptoSetup = &cryptoSetup{} + +// NewCryptoSetupClient creates a new crypto setup for the client +// chromeParrot makes the client emit Chrome's TLS ClientHello via uTLS instead of +// crypto/tls. It forces enable0RTT off: see newUTLSQUICClient for why resumption +// can't be carried across the two TLS stacks. +func NewCryptoSetupClient( + connID protocol.ConnectionID, + tp *wire.TransportParameters, + tlsConf *tls.Config, + enable0RTT bool, + chromeParrot bool, + rttStats *utils.RTTStats, + qlogger qlogwriter.Recorder, + logger utils.Logger, + version protocol.Version, +) (CryptoSetup, error) { + cs := newCryptoSetup( + connID, + tp, + rttStats, + qlogger, + logger, + protocol.PerspectiveClient, + version, + ) + + tlsConf = setupConfigForClient(tlsConf) + cs.tlsConf = tlsConf + cs.allow0RTT = enable0RTT && !chromeParrot + + if chromeParrot { + conn, err := newUTLSQUICClient(tlsConf) + if err != nil { + return nil, err + } + cs.conn = conn + } else { + cs.conn = tls.QUICClient(&tls.QUICConfig{ + TLSConfig: tlsConf, + EnableSessionEvents: true, + }) + } + cs.conn.SetTransportParameters(cs.ourParams.Marshal(protocol.PerspectiveClient)) + + return cs, nil +} + +// NewCryptoSetupServer creates a new crypto setup for the server +func NewCryptoSetupServer( + connID protocol.ConnectionID, + localAddr, remoteAddr net.Addr, + tp *wire.TransportParameters, + tlsConf *tls.Config, + allow0RTT bool, + rttStats *utils.RTTStats, + qlogger qlogwriter.Recorder, + logger utils.Logger, + version protocol.Version, +) CryptoSetup { + cs := newCryptoSetup( + connID, + tp, + rttStats, + qlogger, + logger, + protocol.PerspectiveServer, + version, + ) + cs.allow0RTT = allow0RTT + + tlsConf = setupConfigForServer(tlsConf, localAddr, remoteAddr) + + cs.tlsConf = tlsConf + cs.conn = tls.QUICServer(getQUICConfig(tlsConf, localAddr, remoteAddr)) + return cs +} + +func newCryptoSetup( + connID protocol.ConnectionID, + tp *wire.TransportParameters, + rttStats *utils.RTTStats, + qlogger qlogwriter.Recorder, + logger utils.Logger, + perspective protocol.Perspective, + version protocol.Version, +) *cryptoSetup { + initialSealer, initialOpener := NewInitialAEAD(connID, perspective, version) + if qlogger != nil { + qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.EncryptionInitial, protocol.PerspectiveClient), + }) + qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.EncryptionInitial, protocol.PerspectiveServer), + }) + } + return &cryptoSetup{ + initialSealer: initialSealer, + initialOpener: initialOpener, + aead: newUpdatableAEAD(rttStats, qlogger, logger, version), + events: make([]Event, 0, 16), + ourParams: tp, + rttStats: rttStats, + qlogger: qlogger, + logger: logger, + perspective: perspective, + version: version, + } +} + +func (h *cryptoSetup) ChangeConnectionID(id protocol.ConnectionID) { + initialSealer, initialOpener := NewInitialAEAD(id, h.perspective, h.version) + h.initialSealer = initialSealer + h.initialOpener = initialOpener + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.EncryptionInitial, protocol.PerspectiveClient), + }) + h.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.EncryptionInitial, protocol.PerspectiveServer), + }) + } +} + +func (h *cryptoSetup) SetLargest1RTTAcked(pn protocol.PacketNumber) error { + return h.aead.SetLargestAcked(pn) +} + +func (h *cryptoSetup) StartHandshake(ctx context.Context) error { + err := h.conn.Start(context.WithValue(ctx, QUICVersionContextKey, h.version)) + if err != nil { + return wrapError(err) + } + for { + ev := h.conn.NextEvent() + if err := h.handleEvent(ev); err != nil { + return wrapError(err) + } + if ev.Kind == tls.QUICNoEvent { + break + } + } + if h.perspective == protocol.PerspectiveClient { + if h.zeroRTTSealer != nil && h.zeroRTTParameters != nil { + h.logger.Debugf("Doing 0-RTT.") + h.events = append(h.events, Event{Kind: EventRestoredTransportParameters, TransportParameters: h.zeroRTTParameters}) + } else { + h.logger.Debugf("Not doing 0-RTT. Has sealer: %t, has params: %t", h.zeroRTTSealer != nil, h.zeroRTTParameters != nil) + } + } + return nil +} + +// Close closes the crypto setup. +// It aborts the handshake, if it is still running. +func (h *cryptoSetup) Close() error { + return h.conn.Close() +} + +// HandleMessage handles a TLS handshake message. +// It is called by the crypto streams when a new message is available. +func (h *cryptoSetup) HandleMessage(data []byte, encLevel protocol.EncryptionLevel) error { + if err := h.handleMessage(data, encLevel); err != nil { + return wrapError(err) + } + return nil +} + +func (h *cryptoSetup) handleMessage(data []byte, encLevel protocol.EncryptionLevel) error { + if err := h.conn.HandleData(encLevel.ToTLSEncryptionLevel(), data); err != nil { + return err + } + for { + ev := h.conn.NextEvent() + if err := h.handleEvent(ev); err != nil { + return err + } + if ev.Kind == tls.QUICNoEvent { + return nil + } + } +} + +func (h *cryptoSetup) handleEvent(ev tls.QUICEvent) (err error) { + switch ev.Kind { + case tls.QUICNoEvent: + return nil + case tls.QUICSetReadSecret: + h.setReadKey(ev.Level, ev.Suite, ev.Data) + return nil + case tls.QUICSetWriteSecret: + h.setWriteKey(ev.Level, ev.Suite, ev.Data) + return nil + case tls.QUICTransportParameters: + return h.handleTransportParameters(ev.Data) + case tls.QUICTransportParametersRequired: + h.conn.SetTransportParameters(h.ourParams.Marshal(h.perspective)) + return nil + case tls.QUICRejectedEarlyData: + h.rejected0RTT() + return nil + case tls.QUICWriteData: + h.writeRecord(ev.Level, ev.Data) + return nil + case tls.QUICHandshakeDone: + h.handshakeComplete() + return nil + case tls.QUICStoreSession: + if h.perspective == protocol.PerspectiveServer { + panic("cryptoSetup BUG: unexpected QUICStoreSession event for the server") + } + ev.SessionState.Extra = append( + ev.SessionState.Extra, + addSessionStateExtraPrefix(h.marshalDataForSessionState(ev.SessionState.EarlyData)), + ) + return h.conn.StoreSession(ev.SessionState) + case tls.QUICResumeSession: + var allowEarlyData bool + switch h.perspective { + case protocol.PerspectiveClient: + // for clients, this event occurs when a session ticket is selected + allowEarlyData = h.handleDataFromSessionState( + findSessionStateExtraData(ev.SessionState.Extra), + ev.SessionState.EarlyData, + ) + case protocol.PerspectiveServer: + // for servers, this event occurs when receiving the client's session ticket + allowEarlyData = h.handleSessionTicket( + findSessionStateExtraData(ev.SessionState.Extra), + ev.SessionState.EarlyData, + ) + } + if ev.SessionState.EarlyData { + ev.SessionState.EarlyData = allowEarlyData + } + return nil + case quicErrorEvent: + return extractQUICEventError(ev) + default: + // Unknown events should be ignored. + // crypto/tls will ensure that this is safe to do. + // See the discussion following https://github.com/golang/go/issues/68124#issuecomment-2187042510 for details. + return nil + } +} + +func (h *cryptoSetup) NextEvent() Event { + if len(h.events) == 0 { + return Event{Kind: EventNoEvent} + } + ev := h.events[0] + h.events = h.events[1:] + return ev +} + +func (h *cryptoSetup) handleTransportParameters(data []byte) error { + var tp wire.TransportParameters + if err := tp.Unmarshal(data, h.perspective.Opposite()); err != nil { + return err + } + h.peerParams = &tp + h.events = append(h.events, Event{Kind: EventReceivedTransportParameters, TransportParameters: h.peerParams}) + return nil +} + +// must be called after receiving the transport parameters +func (h *cryptoSetup) marshalDataForSessionState(earlyData bool) []byte { + b := make([]byte, 0, 256) + b = quicvarint.Append(b, clientSessionStateRevision) + if earlyData { + // only save the transport parameters for 0-RTT enabled session tickets + return h.peerParams.MarshalForSessionTicket(b) + } + return b +} + +func (h *cryptoSetup) handleDataFromSessionState(data []byte, earlyData bool) (allowEarlyData bool) { + tp, err := decodeDataFromSessionState(data, earlyData) + if err != nil { + h.logger.Debugf("Restoring of transport parameters from session ticket failed: %s", err.Error()) + return + } + // The session ticket might have been saved from a connection that allowed 0-RTT, + // and therefore contain transport parameters. + // Only use them if 0-RTT is actually used on the new connection. + if tp != nil && h.allow0RTT { + h.zeroRTTParameters = tp + return true + } + return false +} + +func decodeDataFromSessionState(b []byte, earlyData bool) (*wire.TransportParameters, error) { + ver, l, err := quicvarint.Parse(b) + if err != nil { + return nil, err + } + b = b[l:] + if ver != clientSessionStateRevision { + return nil, fmt.Errorf("mismatching version. Got %d, expected %d", ver, clientSessionStateRevision) + } + if !earlyData { + return nil, nil + } + var tp wire.TransportParameters + if err := tp.UnmarshalFromSessionTicket(b); err != nil { + return nil, err + } + return &tp, nil +} + +func (h *cryptoSetup) getDataForSessionTicket() []byte { + return (&sessionTicket{ + Parameters: h.ourParams, + }).Marshal() +} + +// GetSessionTicket generates a new session ticket. +// Due to limitations in crypto/tls, it's only possible to generate a single session ticket per connection. +// It is only valid for the server. +func (h *cryptoSetup) GetSessionTicket() ([]byte, error) { + if err := h.conn.SendSessionTicket(tls.QUICSessionTicketOptions{ + EarlyData: h.allow0RTT, + Extra: [][]byte{addSessionStateExtraPrefix(h.getDataForSessionTicket())}, + }); err != nil { + // Session tickets might be disabled by tls.Config.SessionTicketsDisabled. + // We can't check h.tlsConfig here, since the actual config might have been obtained from + // the GetConfigForClient callback. + // See https://github.com/golang/go/issues/62032. + // This error assertion can be removed once we drop support for Go 1.25. + if strings.Contains(err.Error(), "session ticket keys unavailable") { + return nil, nil + } + return nil, err + } + // If session tickets are disabled, NextEvent will immediately return QUICNoEvent, + // and we will return a nil ticket. + var ticket []byte + for { + ev := h.conn.NextEvent() + if ev.Kind == tls.QUICNoEvent { + break + } + if ev.Kind == tls.QUICWriteData && ev.Level == tls.QUICEncryptionLevelApplication { + if ticket != nil { + h.logger.Errorf("unexpected multiple session tickets") + continue + } + ticket = ev.Data + } else { + h.logger.Errorf("unexpected event: %v", ev.Kind) + } + } + return ticket, nil +} + +// handleSessionTicket is called for the server when receiving the client's session ticket. +// It reads parameters from the session ticket and checks whether to accept 0-RTT if the session ticket enabled 0-RTT. +// Note that the fact that the session ticket allows 0-RTT doesn't mean that the actual TLS handshake enables 0-RTT: +// A client may use a 0-RTT enabled session to resume a TLS session without using 0-RTT. +func (h *cryptoSetup) handleSessionTicket(data []byte, using0RTT bool) (allowEarlyData bool) { + var t sessionTicket + if err := t.Unmarshal(data); err != nil { + h.logger.Debugf("Unmarshalling session ticket failed: %s", err.Error()) + return false + } + if !using0RTT { + return false + } + valid := h.ourParams.ValidFor0RTT(t.Parameters) + if !valid { + h.logger.Debugf("Transport parameters changed. Rejecting 0-RTT.") + return false + } + if !h.allow0RTT { + h.logger.Debugf("0-RTT not allowed. Rejecting 0-RTT.") + return false + } + return true +} + +// rejected0RTT is called for the client when the server rejects 0-RTT. +func (h *cryptoSetup) rejected0RTT() { + h.logger.Debugf("0-RTT was rejected. Dropping 0-RTT keys.") + + had0RTTKeys := h.zeroRTTSealer != nil + h.zeroRTTSealer = nil + + if had0RTTKeys { + h.events = append(h.events, Event{Kind: EventDiscard0RTTKeys}) + } +} + +func (h *cryptoSetup) setReadKey(el tls.QUICEncryptionLevel, suiteID uint16, trafficSecret []byte) { + suite := getCipherSuite(suiteID) + //nolint:exhaustive // The TLS stack doesn't export Initial keys. + switch el { + case tls.QUICEncryptionLevelEarly: + if h.perspective == protocol.PerspectiveClient { + panic("Received 0-RTT read key for the client") + } + h.zeroRTTOpener = newLongHeaderOpener( + createAEAD(suite, trafficSecret, h.version), + newHeaderProtector(suite, trafficSecret, true, h.version), + ) + h.used0RTT.Store(true) + if h.logger.Debug() { + h.logger.Debugf("Installed 0-RTT Read keys (using %s)", tls.CipherSuiteName(suite.ID)) + } + case tls.QUICEncryptionLevelHandshake: + h.handshakeOpener = newLongHeaderOpener( + createAEAD(suite, trafficSecret, h.version), + newHeaderProtector(suite, trafficSecret, true, h.version), + ) + if h.logger.Debug() { + h.logger.Debugf("Installed Handshake Read keys (using %s)", tls.CipherSuiteName(suite.ID)) + } + case tls.QUICEncryptionLevelApplication: + h.aead.SetReadKey(suite, trafficSecret) + h.has1RTTOpener = true + if h.logger.Debug() { + h.logger.Debugf("Installed 1-RTT Read keys (using %s)", tls.CipherSuiteName(suite.ID)) + } + default: + panic("unexpected read encryption level") + } + h.events = append(h.events, Event{Kind: EventReceivedReadKeys}) + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.FromTLSEncryptionLevel(el), h.perspective.Opposite()), + }) + } +} + +func (h *cryptoSetup) setWriteKey(el tls.QUICEncryptionLevel, suiteID uint16, trafficSecret []byte) { + suite := getCipherSuite(suiteID) + //nolint:exhaustive // The TLS stack doesn't export Initial keys. + switch el { + case tls.QUICEncryptionLevelEarly: + if h.perspective == protocol.PerspectiveServer { + panic("Received 0-RTT write key for the server") + } + h.zeroRTTSealer = newLongHeaderSealer( + createAEAD(suite, trafficSecret, h.version), + newHeaderProtector(suite, trafficSecret, true, h.version), + ) + if h.logger.Debug() { + h.logger.Debugf("Installed 0-RTT Write keys (using %s)", tls.CipherSuiteName(suite.ID)) + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.Encryption0RTT, h.perspective), + }) + } + // don't set used0RTT here. 0-RTT might still get rejected. + return + case tls.QUICEncryptionLevelHandshake: + h.handshakeSealer = newLongHeaderSealer( + createAEAD(suite, trafficSecret, h.version), + newHeaderProtector(suite, trafficSecret, true, h.version), + ) + if h.logger.Debug() { + h.logger.Debugf("Installed Handshake Write keys (using %s)", tls.CipherSuiteName(suite.ID)) + } + case tls.QUICEncryptionLevelApplication: + h.aead.SetWriteKey(suite, trafficSecret) + h.has1RTTSealer = true + if h.logger.Debug() { + h.logger.Debugf("Installed 1-RTT Write keys (using %s)", tls.CipherSuiteName(suite.ID)) + } + if h.zeroRTTSealer != nil { + // Once we receive handshake keys, we know that 0-RTT was not rejected. + h.used0RTT.Store(true) + h.zeroRTTSealer = nil + h.logger.Debugf("Dropping 0-RTT keys.") + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyDiscarded{KeyType: qlog.KeyTypeClient0RTT}) + } + } + default: + panic("unexpected write encryption level") + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateTLS, + KeyType: encLevelToKeyType(protocol.FromTLSEncryptionLevel(el), h.perspective), + }) + } +} + +// writeRecord is called when TLS writes data +func (h *cryptoSetup) writeRecord(encLevel tls.QUICEncryptionLevel, p []byte) { + //nolint:exhaustive // handshake records can only be written for Initial and Handshake. + switch encLevel { + case tls.QUICEncryptionLevelInitial: + h.events = append(h.events, Event{Kind: EventWriteInitialData, Data: p}) + case tls.QUICEncryptionLevelHandshake: + h.events = append(h.events, Event{Kind: EventWriteHandshakeData, Data: p}) + case tls.QUICEncryptionLevelApplication: + panic("unexpected write") + default: + panic(fmt.Sprintf("unexpected write encryption level: %s", encLevel)) + } +} + +func (h *cryptoSetup) DiscardInitialKeys() { + dropped := h.initialOpener != nil + h.initialOpener = nil + h.initialSealer = nil + if dropped { + h.logger.Debugf("Dropping Initial keys.") + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyDiscarded{KeyType: qlog.KeyTypeClientInitial}) + h.qlogger.RecordEvent(qlog.KeyDiscarded{KeyType: qlog.KeyTypeServerInitial}) + } + } +} + +func (h *cryptoSetup) handshakeComplete() { + h.handshakeCompleteTime = time.Now() + h.events = append(h.events, Event{Kind: EventHandshakeComplete}) +} + +func (h *cryptoSetup) SetHandshakeConfirmed() { + h.aead.SetHandshakeConfirmed() + // drop Handshake keys + var dropped bool + if h.handshakeOpener != nil { + h.handshakeOpener = nil + h.handshakeSealer = nil + dropped = true + } + if dropped { + h.logger.Debugf("Dropping Handshake keys.") + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyDiscarded{KeyType: qlog.KeyTypeClientHandshake}) + h.qlogger.RecordEvent(qlog.KeyDiscarded{KeyType: qlog.KeyTypeServerHandshake}) + } + } +} + +func (h *cryptoSetup) GetInitialSealer() (LongHeaderSealer, error) { + if h.initialSealer == nil { + return nil, ErrKeysDropped + } + return h.initialSealer, nil +} + +func (h *cryptoSetup) Get0RTTSealer() (LongHeaderSealer, error) { + if h.zeroRTTSealer == nil { + return nil, ErrKeysDropped + } + return h.zeroRTTSealer, nil +} + +func (h *cryptoSetup) GetHandshakeSealer() (LongHeaderSealer, error) { + if h.handshakeSealer == nil { + if h.initialSealer == nil { + return nil, ErrKeysDropped + } + return nil, ErrKeysNotYetAvailable + } + return h.handshakeSealer, nil +} + +func (h *cryptoSetup) Get1RTTSealer() (ShortHeaderSealer, error) { + if !h.has1RTTSealer { + return nil, ErrKeysNotYetAvailable + } + return h.aead, nil +} + +func (h *cryptoSetup) GetInitialOpener() (LongHeaderOpener, error) { + if h.initialOpener == nil { + return nil, ErrKeysDropped + } + return h.initialOpener, nil +} + +func (h *cryptoSetup) Get0RTTOpener() (LongHeaderOpener, error) { + if h.zeroRTTOpener == nil { + if h.initialOpener != nil { + return nil, ErrKeysNotYetAvailable + } + // if the initial opener is also not available, the keys were already dropped + return nil, ErrKeysDropped + } + return h.zeroRTTOpener, nil +} + +func (h *cryptoSetup) GetHandshakeOpener() (LongHeaderOpener, error) { + if h.handshakeOpener == nil { + if h.initialOpener != nil { + return nil, ErrKeysNotYetAvailable + } + // if the initial opener is also not available, the keys were already dropped + return nil, ErrKeysDropped + } + return h.handshakeOpener, nil +} + +func (h *cryptoSetup) Get1RTTOpener() (ShortHeaderOpener, error) { + if h.zeroRTTOpener != nil && time.Since(h.handshakeCompleteTime) > 3*h.rttStats.PTO(true) { + h.zeroRTTOpener = nil + h.logger.Debugf("Dropping 0-RTT keys.") + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.KeyDiscarded{KeyType: qlog.KeyTypeClient0RTT}) + } + } + + if !h.has1RTTOpener { + return nil, ErrKeysNotYetAvailable + } + return h.aead, nil +} + +func (h *cryptoSetup) ConnectionState() ConnectionState { + return ConnectionState{ + ConnectionState: h.conn.ConnectionState(), + Used0RTT: h.used0RTT.Load(), + } +} + +func wrapError(err error) error { + if alertErr := tls.AlertError(0); errors.As(err, &alertErr) { + return qerr.NewLocalCryptoError(uint8(alertErr), err) + } + return &qerr.TransportError{ErrorCode: qerr.InternalError, ErrorMessage: err.Error()} +} + +func encLevelToKeyType(encLevel protocol.EncryptionLevel, pers protocol.Perspective) qlog.KeyType { + if pers == protocol.PerspectiveServer { + switch encLevel { + case protocol.EncryptionInitial: + return qlog.KeyTypeServerInitial + case protocol.EncryptionHandshake: + return qlog.KeyTypeServerHandshake + case protocol.Encryption0RTT: + return qlog.KeyTypeServer0RTT + case protocol.Encryption1RTT: + return qlog.KeyTypeServer1RTT + default: + return "" + } + } + switch encLevel { + case protocol.EncryptionInitial: + return qlog.KeyTypeClientInitial + case protocol.EncryptionHandshake: + return qlog.KeyTypeClientHandshake + case protocol.Encryption0RTT: + return qlog.KeyTypeClient0RTT + case protocol.Encryption1RTT: + return qlog.KeyTypeClient1RTT + default: + return "" + } +} diff --git a/third_party/quic-go/internal/handshake/crypto_setup_test.go b/third_party/quic-go/internal/handshake/crypto_setup_test.go new file mode 100644 index 0000000..6ecd6d2 --- /dev/null +++ b/third_party/quic-go/internal/handshake/crypto_setup_test.go @@ -0,0 +1,574 @@ +package handshake + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "errors" + "math/big" + "net" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/testdata" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +const ( + typeClientHello = 1 + typeNewSessionTicket = 4 +) + +type mockClientSessionCache struct { + cache tls.ClientSessionCache + puts chan *tls.ClientSessionState +} + +var _ tls.ClientSessionCache = &mockClientSessionCache{} + +func newMockClientSessionCache() *mockClientSessionCache { + return &mockClientSessionCache{ + puts: make(chan *tls.ClientSessionState, 1), + cache: tls.NewLRUClientSessionCache(1), + } +} + +func (m *mockClientSessionCache) Get(sessionKey string) (session *tls.ClientSessionState, ok bool) { + return m.cache.Get(sessionKey) +} + +func (m *mockClientSessionCache) Put(sessionKey string, cs *tls.ClientSessionState) { + m.puts <- cs + m.cache.Put(sessionKey, cs) +} + +func getTLSConfigs() (clientConf, serverConf *tls.Config) { + clientConf = &tls.Config{ + ServerName: "localhost", + RootCAs: testdata.GetRootCA(), + NextProtos: []string{"crypto-setup"}, + } + serverConf = testdata.GetTLSConfig() + serverConf.NextProtos = []string{"crypto-setup"} + return clientConf, serverConf +} + +func TestErrorBeforeClientHelloGeneration(t *testing.T) { + tlsConf := testdata.GetTLSConfig() + tlsConf.InsecureSkipVerify = true + tlsConf.NextProtos = []string{""} + cl, err := NewCryptoSetupClient( + protocol.ConnectionID{}, + &wire.TransportParameters{}, + tlsConf, + false, + false, + utils.NewRTTStats(), + nil, + utils.DefaultLogger.WithPrefix("client"), + protocol.Version1, + ) + require.NoError(t, err) + + var terr *qerr.TransportError + err = cl.StartHandshake(context.Background()) + require.True(t, errors.As(err, &terr)) + require.Equal(t, uint64(0x100+0x50), uint64(terr.ErrorCode)) + require.Contains(t, err.Error(), "tls: invalid NextProtos value") +} + +func TestMessageReceivedAtWrongEncryptionLevel(t *testing.T) { + var token protocol.StatelessResetToken + server := NewCryptoSetupServer( + protocol.ConnectionID{}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 1234}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 4321}, + &wire.TransportParameters{StatelessResetToken: &token}, + testdata.GetTLSConfig(), + false, + utils.NewRTTStats(), + nil, + utils.DefaultLogger.WithPrefix("server"), + protocol.Version1, + ) + + require.NoError(t, server.StartHandshake(context.Background())) + + fakeCH := append([]byte{typeClientHello, 0, 0, 6}, []byte("foobar")...) + // wrong encryption level + err := server.HandleMessage(fakeCH, protocol.EncryptionHandshake) + require.Error(t, err) + require.Contains(t, err.Error(), "tls: handshake data received at wrong level") +} + +// The clientEvents and serverEvents contain all events that were not processed by the function, +// i.e. not EventWriteInitialData, EventWriteHandshakeData, EventHandshakeComplete. +func handshake(t *testing.T, client, server CryptoSetup) (clientEvents []Event, clientErr error, serverEvents []Event, serverErr error) { + t.Helper() + require.NoError(t, client.StartHandshake(context.Background())) + require.NoError(t, server.StartHandshake(context.Background())) + + var clientHandshakeComplete, serverHandshakeComplete bool + + for { + clientLoop: + for { + ev := client.NextEvent() + switch ev.Kind { + case EventNoEvent: + break clientLoop + case EventWriteInitialData: + serverErr = server.HandleMessage(ev.Data, protocol.EncryptionInitial) + if serverErr != nil { + return + } + case EventWriteHandshakeData: + serverErr = server.HandleMessage(ev.Data, protocol.EncryptionHandshake) + if serverErr != nil { + return + } + case EventHandshakeComplete: + clientHandshakeComplete = true + default: + clientEvents = append(clientEvents, ev) + } + } + + serverLoop: + for { + ev := server.NextEvent() + switch ev.Kind { + case EventNoEvent: + break serverLoop + case EventWriteInitialData: + clientErr = client.HandleMessage(ev.Data, protocol.EncryptionInitial) + if clientErr != nil { + return + } + case EventWriteHandshakeData: + clientErr = client.HandleMessage(ev.Data, protocol.EncryptionHandshake) + if clientErr != nil { + return + } + case EventHandshakeComplete: + serverHandshakeComplete = true + ticket, err := server.GetSessionTicket() + require.NoError(t, err) + if ticket != nil { + require.NoError(t, client.HandleMessage(ticket, protocol.Encryption1RTT)) + } + default: + serverEvents = append(serverEvents, ev) + } + } + + if clientHandshakeComplete && serverHandshakeComplete { + break + } + } + return +} + +func handshakeWithTLSConf( + t *testing.T, + clientConf, serverConf *tls.Config, + clientRTTStats, serverRTTStats *utils.RTTStats, + clientTransportParameters, serverTransportParameters *wire.TransportParameters, + enable0RTT bool, +) (CryptoSetup /* client */, []Event /* more client events */, error, /* client error */ + CryptoSetup /* server */, []Event /* more server events */, error, /* server error */ +) { + t.Helper() + client, err := NewCryptoSetupClient( + protocol.ConnectionID{}, + clientTransportParameters, + clientConf, + enable0RTT, + false, + clientRTTStats, + nil, + utils.DefaultLogger.WithPrefix("client"), + protocol.Version1, + ) + require.NoError(t, err) + + if serverTransportParameters.StatelessResetToken == nil { + var token protocol.StatelessResetToken + serverTransportParameters.StatelessResetToken = &token + } + server := NewCryptoSetupServer( + protocol.ConnectionID{}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 1234}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 4321}, + serverTransportParameters, + serverConf, + enable0RTT, + serverRTTStats, + nil, + utils.DefaultLogger.WithPrefix("server"), + protocol.Version1, + ) + cEvents, cErr, sEvents, sErr := handshake(t, client, server) + return client, cEvents, cErr, server, sEvents, sErr +} + +func TestHandshake(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + _, _, clientErr, _, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) +} + +func TestHelloRetryRequest(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + serverConf.CurvePreferences = []tls.CurveID{tls.CurveP384} + _, _, clientErr, _, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) +} + +func TestWithClientAuth(t *testing.T) { + pub, priv, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + tmpl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{}, + SignatureAlgorithm: x509.PureEd25519, + NotBefore: time.Now(), + NotAfter: time.Now().Add(time.Hour), + BasicConstraintsValid: true, + } + certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, pub, priv) + require.NoError(t, err) + clientCert := tls.Certificate{ + PrivateKey: priv, + Certificate: [][]byte{certDER}, + } + + clientConf, serverConf := getTLSConfigs() + clientConf.Certificates = []tls.Certificate{clientCert} + serverConf.ClientAuth = tls.RequireAnyClientCert + _, _, clientErr, _, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) +} + +func TestTransportParameters(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + cTransportParameters := &wire.TransportParameters{ActiveConnectionIDLimit: 2, MaxIdleTimeout: 42 * time.Second} + client, err := NewCryptoSetupClient( + protocol.ConnectionID{}, + cTransportParameters, + clientConf, + false, + false, + utils.NewRTTStats(), + nil, + utils.DefaultLogger.WithPrefix("client"), + protocol.Version1, + ) + require.NoError(t, err) + + var token protocol.StatelessResetToken + sTransportParameters := &wire.TransportParameters{ + MaxIdleTimeout: 1337 * time.Second, + StatelessResetToken: &token, + ActiveConnectionIDLimit: 2, + } + server := NewCryptoSetupServer( + protocol.ConnectionID{}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 1234}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 4321}, + sTransportParameters, + serverConf, + false, + utils.NewRTTStats(), + nil, + utils.DefaultLogger.WithPrefix("server"), + protocol.Version1, + ) + + clientEvents, cErr, serverEvents, sErr := handshake(t, client, server) + require.NoError(t, cErr) + require.NoError(t, sErr) + var clientReceivedTransportParameters *wire.TransportParameters + for _, ev := range clientEvents { + if ev.Kind == EventReceivedTransportParameters { + clientReceivedTransportParameters = ev.TransportParameters + } + } + require.NotNil(t, clientReceivedTransportParameters) + require.Equal(t, 1337*time.Second, clientReceivedTransportParameters.MaxIdleTimeout) + + var serverReceivedTransportParameters *wire.TransportParameters + for _, ev := range serverEvents { + if ev.Kind == EventReceivedTransportParameters { + serverReceivedTransportParameters = ev.TransportParameters + } + } + require.NotNil(t, serverReceivedTransportParameters) + require.Equal(t, 42*time.Second, serverReceivedTransportParameters.MaxIdleTimeout) +} + +func TestNewSessionTicketAtWrongEncryptionLevel(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + client, _, clientErr, _, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + + // inject an invalid session ticket + b := append([]byte{uint8(typeNewSessionTicket), 0, 0, 6}, []byte("foobar")...) + err := client.HandleMessage(b, protocol.EncryptionHandshake) + require.Error(t, err) + require.Contains(t, err.Error(), "tls: handshake data received at wrong level") +} + +func TestHandlingNewSessionTicketFails(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + client, _, clientErr, _, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + + // inject an invalid session ticket + b := append([]byte{uint8(typeNewSessionTicket), 0, 0, 6}, []byte("foobar")...) + err := client.HandleMessage(b, protocol.Encryption1RTT) + require.IsType(t, &qerr.TransportError{}, err) + require.True(t, err.(*qerr.TransportError).ErrorCode.IsCryptoError()) +} + +func TestSessionResumption(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + csc := newMockClientSessionCache() + clientConf.ClientSessionCache = csc + client, _, clientErr, server, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + select { + case <-csc.puts: + case <-time.After(time.Second): + t.Fatal("didn't receive a session ticket") + } + require.False(t, server.ConnectionState().DidResume) + require.False(t, client.ConnectionState().DidResume) + + clientRTTStats := utils.NewRTTStats() + serverRTTStats := utils.NewRTTStats() + client, _, clientErr, server, _, serverErr = handshakeWithTLSConf( + t, + clientConf, serverConf, + clientRTTStats, serverRTTStats, + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + select { + case <-csc.puts: + case <-time.After(time.Second): + t.Fatal("didn't receive a session ticket") + } + require.True(t, server.ConnectionState().DidResume) + require.True(t, client.ConnectionState().DidResume) +} + +func TestSessionResumptionDisabled(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + csc := newMockClientSessionCache() + clientConf.ClientSessionCache = csc + client, _, clientErr, server, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + select { + case <-csc.puts: + case <-time.After(time.Second): + t.Fatal("didn't receive a session ticket") + } + require.False(t, server.ConnectionState().DidResume) + require.False(t, client.ConnectionState().DidResume) + + serverConf.SessionTicketsDisabled = true + client, _, clientErr, server, _, serverErr = handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + false, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + select { + case <-csc.puts: + t.Fatal("didn't expect to receive a session ticket") + case <-time.After(25 * time.Millisecond): + } + require.False(t, server.ConnectionState().DidResume) + require.False(t, client.ConnectionState().DidResume) +} + +func Test0RTT(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + csc := newMockClientSessionCache() + clientConf.ClientSessionCache = csc + const initialMaxData protocol.ByteCount = 1337 + client, _, clientErr, server, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + &wire.TransportParameters{ActiveConnectionIDLimit: 2, InitialMaxData: initialMaxData}, + true, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + select { + case <-csc.puts: + case <-time.After(time.Second): + t.Fatal("didn't receive a session ticket") + } + require.False(t, server.ConnectionState().DidResume) + require.False(t, client.ConnectionState().DidResume) + + client, clientEvents, clientErr, server, serverEvents, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + &wire.TransportParameters{ActiveConnectionIDLimit: 2, InitialMaxData: initialMaxData}, + true, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + + var tp *wire.TransportParameters + var clientReceived0RTTKeys bool + for _, ev := range clientEvents { + switch ev.Kind { + case EventRestoredTransportParameters: + tp = ev.TransportParameters + case EventReceivedReadKeys: + clientReceived0RTTKeys = true + } + } + require.True(t, clientReceived0RTTKeys) + require.NotNil(t, tp) + require.Equal(t, initialMaxData, tp.InitialMaxData) + + var serverReceived0RTTKeys bool + for _, ev := range serverEvents { + switch ev.Kind { + case EventReceivedReadKeys: + serverReceived0RTTKeys = true + } + } + require.True(t, serverReceived0RTTKeys) + + require.True(t, server.ConnectionState().DidResume) + require.True(t, client.ConnectionState().DidResume) + require.True(t, server.ConnectionState().Used0RTT) + require.True(t, client.ConnectionState().Used0RTT) +} + +func Test0RTTRejectionOnTransportParametersChanged(t *testing.T) { + clientConf, serverConf := getTLSConfigs() + csc := newMockClientSessionCache() + clientConf.ClientSessionCache = csc + const initialMaxData protocol.ByteCount = 1337 + client, _, clientErr, server, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + utils.NewRTTStats(), utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + &wire.TransportParameters{ActiveConnectionIDLimit: 2, InitialMaxData: initialMaxData}, + true, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + select { + case <-csc.puts: + case <-time.After(time.Second): + t.Fatal("didn't receive a session ticket") + } + require.False(t, server.ConnectionState().DidResume) + require.False(t, client.ConnectionState().DidResume) + + clientRTTStats := utils.NewRTTStats() + client, clientEvents, clientErr, server, _, serverErr := handshakeWithTLSConf( + t, + clientConf, serverConf, + clientRTTStats, utils.NewRTTStats(), + &wire.TransportParameters{ActiveConnectionIDLimit: 2}, + &wire.TransportParameters{ActiveConnectionIDLimit: 2, InitialMaxData: initialMaxData - 1}, + true, + ) + require.NoError(t, clientErr) + require.NoError(t, serverErr) + + var tp *wire.TransportParameters + var clientReceived0RTTKeys bool + for _, ev := range clientEvents { + switch ev.Kind { + case EventRestoredTransportParameters: + tp = ev.TransportParameters + case EventReceivedReadKeys: + clientReceived0RTTKeys = true + } + } + require.True(t, clientReceived0RTTKeys) + require.NotNil(t, tp) + require.Equal(t, initialMaxData, tp.InitialMaxData) + + require.True(t, server.ConnectionState().DidResume) + require.True(t, client.ConnectionState().DidResume) + require.False(t, server.ConnectionState().Used0RTT) + require.False(t, client.ConnectionState().Used0RTT) +} diff --git a/third_party/quic-go/internal/handshake/fake_conn.go b/third_party/quic-go/internal/handshake/fake_conn.go new file mode 100644 index 0000000..54af823 --- /dev/null +++ b/third_party/quic-go/internal/handshake/fake_conn.go @@ -0,0 +1,21 @@ +package handshake + +import ( + "net" + "time" +) + +type conn struct { + localAddr, remoteAddr net.Addr +} + +var _ net.Conn = &conn{} + +func (c *conn) Read([]byte) (int, error) { return 0, nil } +func (c *conn) Write([]byte) (int, error) { return 0, nil } +func (c *conn) Close() error { return nil } +func (c *conn) RemoteAddr() net.Addr { return c.remoteAddr } +func (c *conn) LocalAddr() net.Addr { return c.localAddr } +func (c *conn) SetReadDeadline(time.Time) error { return nil } +func (c *conn) SetWriteDeadline(time.Time) error { return nil } +func (c *conn) SetDeadline(time.Time) error { return nil } diff --git a/third_party/quic-go/internal/handshake/fips140_go126.go b/third_party/quic-go/internal/handshake/fips140_go126.go new file mode 100644 index 0000000..4241794 --- /dev/null +++ b/third_party/quic-go/internal/handshake/fips140_go126.go @@ -0,0 +1,9 @@ +//go:build go1.26 + +package handshake + +import "crypto/fips140" + +func withoutFIPSEnforcement(f func()) { + fips140.WithoutEnforcement(f) +} diff --git a/third_party/quic-go/internal/handshake/fips140_legacy.go b/third_party/quic-go/internal/handshake/fips140_legacy.go new file mode 100644 index 0000000..bc8e639 --- /dev/null +++ b/third_party/quic-go/internal/handshake/fips140_legacy.go @@ -0,0 +1,7 @@ +//go:build !go1.26 + +package handshake + +func withoutFIPSEnforcement(f func()) { + f() +} diff --git a/third_party/quic-go/internal/handshake/handshake_fuzz_test.go b/third_party/quic-go/internal/handshake/handshake_fuzz_test.go new file mode 100644 index 0000000..210e6c0 --- /dev/null +++ b/third_party/quic-go/internal/handshake/handshake_fuzz_test.go @@ -0,0 +1,399 @@ +package handshake + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "errors" + "math/big" + "net" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qtls" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" +) + +var ( + fuzzCert, fuzzAltCert *tls.Certificate + fuzzCertPool *x509.CertPool +) + +func init() { + var err error + fuzzCert, fuzzCertPool, err = generateFuzzCertificate() + if err != nil { + panic(err) + } + fuzzAltCert, _, err = generateFuzzCertificate() + if err != nil { + panic(err) + } +} + +func generateFuzzCertificate() (*tls.Certificate, *x509.CertPool, error) { + _, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + tmpl := x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{Organization: []string{"quic-go fuzzer"}}, + NotBefore: time.Now().Add(-24 * time.Hour), + NotAfter: time.Now().Add(30 * 24 * time.Hour), + KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + DNSNames: []string{"localhost"}, + BasicConstraintsValid: true, + } + derBytes, err := x509.CreateCertificate(rand.Reader, &tmpl, &tmpl, priv.Public(), priv) + if err != nil { + return nil, nil, err + } + cert, err := x509.ParseCertificate(derBytes) + if err != nil { + return nil, nil, err + } + pool := x509.NewCertPool() + pool.AddCert(cert) + return &tls.Certificate{ + Certificate: [][]byte{derBytes}, + PrivateKey: priv, + }, pool, nil +} + +const fuzzALPN = "fuzzing" + +var fuzzSessionTicketKey = [32]byte{ + 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, + 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, +} + +func FuzzHandshake(f *testing.F) { + corpus := ossfuzzseeds.New(f) + + corpus.Add( + uint8(0), // cipherSuite + uint8(0), // clientAuth + uint8(0xFF), // messageToReplace: won't match + uint8(0), // messageEncLevel + uint8(0), // zeroRTTMode + uint8(0), // postHandshakeTarget + uint8(0), // invalidTP + uint8(0), // sessionMode + uint8(0), // tlsCallbacks + uint8(0), // alpnMode + []byte("foobar"), + ) + corpus.Add( + uint8(2), // cipherSuite: ChaCha20 + uint8(0), // clientAuth + uint8(1), // messageToReplace: ClientHello + uint8(0), // messageEncLevel: Initial + uint8(3), // zeroRTTMode: both + uint8(0), // postHandshakeTarget + uint8(0), // invalidTP + uint8(1), // sessionMode: cache+ticket+enabled + uint8(0), // tlsCallbacks + uint8(0), // alpnMode + []byte("hello world"), + ) + + // Parameters: + // - cipherSuite (0-3): AES-128-GCM, AES-256-GCM, ChaCha20, default + // - clientAuth (0-4): maps to tls.ClientAuthType + // - messageToReplace: TLS message type byte to replace with fuzz data + // - messageEncLevel (0-2): encryption level for replaced message (Initial, Handshake, 1-RTT) + // - zeroRTTMode (0-3): 0=neither, 1=client, 2=server, 3=both + // - postHandshakeTarget (0-3): 0=neither, 1=client, 2=server, 3=both + // - invalidTP (0-3): 0=both valid, 1=client invalid, 2=server invalid, 3=both invalid + // - sessionMode (0-3): 0=no cache, 1=cache+ticket+enabled, 2=cache+disabled, 3=cache+no ticket + // - tlsCallbacks (0-15): getConfigForClient (val%4) x getCertificate (val/4) + // each: 0=off, 1=return value, 2=return error, 3=return nil + // - alpnMode (0-2): 0=correct, 1=wrong, 2=both correct and wrong + f.Fuzz(func(t *testing.T, cipherSuite, clientAuth, messageToReplace, messageEncLevel, zeroRTTMode, postHandshakeTarget, invalidTP, sessionMode, tlsCallbacks, alpnMode uint8, data []byte) { + if len(data) == 0 { + return + } + if cipherSuite > 3 || clientAuth > 4 || messageEncLevel > 2 || zeroRTTMode > 3 || postHandshakeTarget > 3 || invalidTP > 3 || sessionMode > 3 || tlsCallbacks > 15 || alpnMode > 2 { + return + } + + clientConf := &tls.Config{ + MinVersion: tls.VersionTLS13, + ServerName: "localhost", + NextProtos: []string{fuzzALPN}, + RootCAs: fuzzCertPool, + } + if sessionMode > 0 { + clientConf.ClientSessionCache = tls.NewLRUClientSessionCache(5) + } + + fuzzRunHandshake(t, cipherSuite, clientAuth, messageToReplace, messageEncLevel, zeroRTTMode, postHandshakeTarget, invalidTP, sessionMode, tlsCallbacks, alpnMode, clientConf, data) + fuzzRunHandshake(t, cipherSuite, clientAuth, messageToReplace, messageEncLevel, zeroRTTMode, postHandshakeTarget, invalidTP, sessionMode, tlsCallbacks, alpnMode, clientConf, data) + }) +} + +func fuzzRunHandshake( + t *testing.T, + cipherSuiteVal, clientAuthVal, messageToReplace, messageEncLevelVal, zeroRTTMode, postHandshakeTarget, invalidTP, sessionMode, tlsCallbacks, alpnMode uint8, + clientConf *tls.Config, + data []byte, +) { + t.Helper() + + switch cipherSuiteVal { + case 0: + defer qtls.SetCipherSuite(tls.TLS_AES_128_GCM_SHA256)() + case 1: + defer qtls.SetCipherSuite(tls.TLS_AES_256_GCM_SHA384)() + case 2: + defer qtls.SetCipherSuite(tls.TLS_CHACHA20_POLY1305_SHA256)() + case 3: + // use default cipher suites + } + + var tlsClientAuth tls.ClientAuthType + switch clientAuthVal { + case 0: + tlsClientAuth = tls.NoClientCert + case 1: + tlsClientAuth = tls.RequestClientCert + case 2: + tlsClientAuth = tls.RequireAnyClientCert + case 3: + tlsClientAuth = tls.VerifyClientCertIfGiven + case 4: + tlsClientAuth = tls.RequireAndVerifyClientCert + } + + var msgEncLevel protocol.EncryptionLevel + switch messageEncLevelVal { + case 0: + msgEncLevel = protocol.EncryptionInitial + case 1: + msgEncLevel = protocol.EncryptionHandshake + case 2: + msgEncLevel = protocol.Encryption1RTT + } + + serverConf := &tls.Config{ + MinVersion: tls.VersionTLS13, + Certificates: []tls.Certificate{*fuzzCert}, + NextProtos: []string{fuzzALPN}, + SessionTicketKey: fuzzSessionTicketKey, + ClientAuth: tlsClientAuth, + } + + enable0RTTClient := zeroRTTMode == 1 || zeroRTTMode == 3 + enable0RTTServer := zeroRTTMode == 2 || zeroRTTMode == 3 + sendPostHandshakeToClient := postHandshakeTarget == 1 || postHandshakeTarget == 3 + sendPostHandshakeToServer := postHandshakeTarget == 2 || postHandshakeTarget == 3 + + sendSessionTicket := sessionMode == 1 + serverConf.SessionTicketsDisabled = sessionMode == 2 + + getConfigForClient := tlsCallbacks % 4 + getCertificate := tlsCallbacks / 4 + + if getConfigForClient > 0 { + serverConf.GetConfigForClient = func(*tls.ClientHelloInfo) (*tls.Config, error) { + switch getConfigForClient { + case 1: + return serverConf, nil + case 2: + return nil, errors.New("getting client config failed") + case 3: + return nil, nil + } + return nil, nil + } + } + if getCertificate > 0 { + serverConf.GetCertificate = func(*tls.ClientHelloInfo) (*tls.Certificate, error) { + switch getCertificate { + case 1: + return fuzzAltCert, nil + case 2: + return nil, errors.New("getting certificate failed") + case 3: + return nil, nil + } + return nil, nil + } + } + + switch alpnMode { + case 1: + serverConf.NextProtos = []string{"wrong"} + case 2: + serverConf.NextProtos = []string{"wrong", fuzzALPN} + } + + clientTP := &wire.TransportParameters{ + ActiveConnectionIDLimit: 2, + InitialMaxData: 1 << 20, + InitialMaxStreamDataBidiLocal: 1 << 16, + InitialMaxStreamDataBidiRemote: 1 << 16, + InitialMaxStreamDataUni: 1 << 16, + } + serverTP := &wire.TransportParameters{ + ActiveConnectionIDLimit: 2, + InitialMaxData: 1 << 20, + InitialMaxStreamDataBidiLocal: 1 << 16, + InitialMaxStreamDataBidiRemote: 1 << 16, + InitialMaxStreamDataUni: 1 << 16, + } + if invalidTP == 1 || invalidTP == 3 { + clientTP.MaxAckDelay = protocol.MaxMaxAckDelay + 5 + } + if invalidTP == 2 || invalidTP == 3 { + serverTP.MaxAckDelay = protocol.MaxMaxAckDelay + 5 + } + + client, cerr := NewCryptoSetupClient( + protocol.ConnectionID{}, + clientTP, + clientConf, + enable0RTTClient, + false, + &utils.RTTStats{}, + nil, + utils.DefaultLogger.WithPrefix("client"), + protocol.Version1, + ) + if cerr != nil { + t.Fatal(cerr) + } + if err := client.StartHandshake(context.Background()); err != nil { + t.Fatal(err) + } + defer client.Close() + + server := NewCryptoSetupServer( + protocol.ConnectionID{}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 1234}, + &net.UDPAddr{IP: net.IPv6loopback, Port: 4321}, + serverTP, + serverConf, + enable0RTTServer, + &utils.RTTStats{}, + nil, + utils.DefaultLogger.WithPrefix("server"), + protocol.Version1, + ) + if err := server.StartHandshake(context.Background()); err != nil { + t.Fatal(err) + } + defer server.Close() + + var clientHandshakeComplete, serverHandshakeComplete bool + for { + var processedEvent bool + clientLoop: + for { + ev := client.NextEvent() + switch ev.Kind { + case EventNoEvent: + if !processedEvent && !clientHandshakeComplete { + return + } + break clientLoop + case EventWriteInitialData, EventWriteHandshakeData: + msg := ev.Data + encLevel := protocol.EncryptionInitial + if ev.Kind == EventWriteHandshakeData { + encLevel = protocol.EncryptionHandshake + } + if msg[0] == messageToReplace { + msg = data + encLevel = msgEncLevel + } + if err := server.HandleMessage(msg, encLevel); err != nil { + return + } + case EventHandshakeComplete: + clientHandshakeComplete = true + } + processedEvent = true + } + processedEvent = false + serverLoop: + for { + ev := server.NextEvent() + switch ev.Kind { + case EventNoEvent: + if !processedEvent && !serverHandshakeComplete { + return + } + break serverLoop + case EventWriteInitialData, EventWriteHandshakeData: + msg := ev.Data + encLevel := protocol.EncryptionInitial + if ev.Kind == EventWriteHandshakeData { + encLevel = protocol.EncryptionHandshake + } + if msg[0] == messageToReplace { + msg = data + encLevel = msgEncLevel + } + if err := client.HandleMessage(msg, encLevel); err != nil { + return + } + case EventHandshakeComplete: + serverHandshakeComplete = true + } + processedEvent = true + } + + if serverHandshakeComplete && clientHandshakeComplete { + break + } + } + + _ = client.ConnectionState() + _ = server.ConnectionState() + + sealer, err := client.Get1RTTSealer() + if err != nil { + t.Fatal("expected to get a 1-RTT sealer") + } + opener, err := server.Get1RTTOpener() + if err != nil { + t.Fatal("expected to get a 1-RTT opener") + } + const plaintext = "Lorem ipsum dolor sit amet, consectetur adipiscing elit." + encrypted := sealer.Seal(nil, []byte(plaintext), 1337, []byte("foobar")) + decrypted, err := opener.Open(nil, encrypted, 0, 1337, protocol.KeyPhaseZero, []byte("foobar")) + if err != nil { + t.Fatalf("decrypting message failed: %s", err) + } + if string(decrypted) != plaintext { + t.Fatal("decrypted message doesn't match") + } + + if sendSessionTicket && !serverConf.SessionTicketsDisabled { + ticket, err := server.GetSessionTicket() + if err != nil { + t.Fatalf("error getting session ticket: %s", err) + } + if ticket == nil { + t.Fatal("expected non-nil session ticket") + } + client.HandleMessage(ticket, protocol.Encryption1RTT) + } + + if sendPostHandshakeToClient { + client.HandleMessage(data, msgEncLevel) + } + if sendPostHandshakeToServer { + server.HandleMessage(data, msgEncLevel) + } +} diff --git a/third_party/quic-go/internal/handshake/handshake_helpers_test.go b/third_party/quic-go/internal/handshake/handshake_helpers_test.go new file mode 100644 index 0000000..9949578 --- /dev/null +++ b/third_party/quic-go/internal/handshake/handshake_helpers_test.go @@ -0,0 +1,41 @@ +package handshake + +import ( + "crypto/fips140" + "crypto/tls" + "encoding/hex" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func splitHexString(t *testing.T, s string) (slice []byte) { + t.Helper() + for ss := range strings.SplitSeq(s, " ") { + if ss[0:2] == "0x" { + ss = ss[2:] + } + d, err := hex.DecodeString(ss) + require.NoError(t, err) + slice = append(slice, d...) + } + return +} + +func TestSplitHexString(t *testing.T) { + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, splitHexString(t, "0xdeadbeef")) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, splitHexString(t, "deadbeef")) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, splitHexString(t, "dead beef")) +} + +var cipherSuites = []cipherSuite{ + getCipherSuite(tls.TLS_AES_128_GCM_SHA256), + getCipherSuite(tls.TLS_AES_256_GCM_SHA384), +} + +func init() { + if !fips140.Enabled() { + cipherSuites = append(cipherSuites, getCipherSuite(tls.TLS_CHACHA20_POLY1305_SHA256)) + } +} diff --git a/third_party/quic-go/internal/handshake/header_protector.go b/third_party/quic-go/internal/handshake/header_protector.go new file mode 100644 index 0000000..53f782e --- /dev/null +++ b/third_party/quic-go/internal/handshake/header_protector.go @@ -0,0 +1,134 @@ +package handshake + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/tls" + "encoding/binary" + "fmt" + + "golang.org/x/crypto/chacha20" + + "github.com/apernet/quic-go/internal/protocol" +) + +type headerProtector interface { + EncryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) + DecryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) +} + +func hkdfHeaderProtectionLabel(v protocol.Version) string { + if v == protocol.Version2 { + return "quicv2 hp" + } + return "quic hp" +} + +func newHeaderProtector(suite cipherSuite, trafficSecret []byte, isLongHeader bool, v protocol.Version) headerProtector { + hkdfLabel := hkdfHeaderProtectionLabel(v) + switch suite.ID { + case tls.TLS_AES_128_GCM_SHA256, tls.TLS_AES_256_GCM_SHA384: + return newAESHeaderProtector(suite, trafficSecret, isLongHeader, hkdfLabel) + case tls.TLS_CHACHA20_POLY1305_SHA256: + return newChaChaHeaderProtector(suite, trafficSecret, isLongHeader, hkdfLabel) + default: + panic(fmt.Sprintf("Invalid cipher suite id: %d", suite.ID)) + } +} + +type aesHeaderProtector struct { + mask [16]byte // AES always has a 16 byte block size + block cipher.Block + isLongHeader bool +} + +var _ headerProtector = &aesHeaderProtector{} + +func newAESHeaderProtector(suite cipherSuite, trafficSecret []byte, isLongHeader bool, hkdfLabel string) headerProtector { + hpKey := hkdfExpandLabel(suite.Hash, trafficSecret, []byte{}, hkdfLabel, suite.KeyLen) + block, err := aes.NewCipher(hpKey) + if err != nil { + panic(fmt.Sprintf("error creating new AES cipher: %s", err)) + } + return &aesHeaderProtector{ + block: block, + isLongHeader: isLongHeader, + } +} + +func (p *aesHeaderProtector) DecryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) { + p.apply(sample, firstByte, hdrBytes) +} + +func (p *aesHeaderProtector) EncryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) { + p.apply(sample, firstByte, hdrBytes) +} + +func (p *aesHeaderProtector) apply(sample []byte, firstByte *byte, hdrBytes []byte) { + if len(sample) != len(p.mask) { + panic("invalid sample size") + } + p.block.Encrypt(p.mask[:], sample) + if p.isLongHeader { + *firstByte ^= p.mask[0] & 0xf + } else { + *firstByte ^= p.mask[0] & 0x1f + } + for i := range hdrBytes { + hdrBytes[i] ^= p.mask[i+1] + } +} + +type chachaHeaderProtector struct { + mask [5]byte + + key [32]byte + isLongHeader bool +} + +var _ headerProtector = &chachaHeaderProtector{} + +func newChaChaHeaderProtector(suite cipherSuite, trafficSecret []byte, isLongHeader bool, hkdfLabel string) headerProtector { + hpKey := hkdfExpandLabel(suite.Hash, trafficSecret, []byte{}, hkdfLabel, suite.KeyLen) + + p := &chachaHeaderProtector{ + isLongHeader: isLongHeader, + } + copy(p.key[:], hpKey) + return p +} + +func (p *chachaHeaderProtector) DecryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) { + p.apply(sample, firstByte, hdrBytes) +} + +func (p *chachaHeaderProtector) EncryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) { + p.apply(sample, firstByte, hdrBytes) +} + +func (p *chachaHeaderProtector) apply(sample []byte, firstByte *byte, hdrBytes []byte) { + if len(sample) != 16 { + panic("invalid sample size") + } + for i := range 5 { + p.mask[i] = 0 + } + cipher, err := chacha20.NewUnauthenticatedCipher(p.key[:], sample[4:]) + if err != nil { + panic(err) + } + cipher.SetCounter(binary.LittleEndian.Uint32(sample[:4])) + cipher.XORKeyStream(p.mask[:], p.mask[:]) + p.applyMask(firstByte, hdrBytes) +} + +func (p *chachaHeaderProtector) applyMask(firstByte *byte, hdrBytes []byte) { + if p.isLongHeader { + *firstByte ^= p.mask[0] & 0xf + } else { + *firstByte ^= p.mask[0] & 0x1f + } + for i := range hdrBytes { + hdrBytes[i] ^= p.mask[i+1] + } +} diff --git a/third_party/quic-go/internal/handshake/hkdf.go b/third_party/quic-go/internal/handshake/hkdf.go new file mode 100644 index 0000000..b28c1f1 --- /dev/null +++ b/third_party/quic-go/internal/handshake/hkdf.go @@ -0,0 +1,26 @@ +package handshake + +import ( + "crypto" + "crypto/hkdf" + "encoding/binary" + "fmt" +) + +// hkdfExpandLabel HKDF expands a label as defined in RFC 8446, section 7.1. +func hkdfExpandLabel(hash crypto.Hash, secret, context []byte, label string, length int) []byte { + b := make([]byte, 3, 3+6+len(label)+1+len(context)) + binary.BigEndian.PutUint16(b, uint16(length)) + b[2] = uint8(6 + len(label)) + b = append(b, []byte("tls13 ")...) + b = append(b, []byte(label)...) + b = b[:3+6+len(label)+1] + b[3+6+len(label)] = uint8(len(context)) + b = append(b, context...) + + expanded, err := hkdf.Expand(hash.New, secret, string(b), length) + if err != nil { + panic(fmt.Errorf("quic: HKDF-Expand-Label invocation failed unexpectedly: %v", err)) + } + return expanded +} diff --git a/third_party/quic-go/internal/handshake/hkdf_test.go b/third_party/quic-go/internal/handshake/hkdf_test.go new file mode 100644 index 0000000..e43084f --- /dev/null +++ b/third_party/quic-go/internal/handshake/hkdf_test.go @@ -0,0 +1,76 @@ +package handshake + +import ( + "crypto" + "crypto/cipher" + "crypto/rand" + "crypto/tls" + "testing" + "unsafe" + + "github.com/stretchr/testify/require" +) + +var tls13CipherSuites = []uint16{tls.TLS_AES_128_GCM_SHA256, tls.TLS_AES_256_GCM_SHA384, tls.TLS_CHACHA20_POLY1305_SHA256} + +type cipherSuiteTLS13 struct { + ID uint16 + KeyLen int + AEAD func(key, fixedNonce []byte) cipher.AEAD + Hash crypto.Hash +} + +//go:linkname cipherSuitesTLS13 crypto/tls.cipherSuitesTLS13 +var cipherSuitesTLS13 []unsafe.Pointer + +func cipherSuiteTLS13ByID(id uint16) *cipherSuiteTLS13 { + for _, v := range cipherSuitesTLS13 { + cs := (*cipherSuiteTLS13)(v) + if cs.ID == id { + return cs + } + } + return nil +} + +//go:linkname nextTrafficSecret crypto/tls.(*cipherSuiteTLS13).nextTrafficSecret +func nextTrafficSecret(cs *cipherSuiteTLS13, trafficSecret []byte) []byte + +func TestHKDF(t *testing.T) { + for _, id := range tls13CipherSuites { + t.Run(tls.CipherSuiteName(id), func(t *testing.T) { + cs := cipherSuiteTLS13ByID(id) + expected := nextTrafficSecret(cs, []byte("foobar")) + expanded := hkdfExpandLabel(cs.Hash, []byte("foobar"), nil, "traffic upd", cs.Hash.Size()) + require.Equal(t, expected, expanded) + }) + } +} + +// As of Go 1.24, the standard library and our implementation of hkdfExpandLabel should provide the same performance. +func BenchmarkHKDFExpandLabelStandardLibrary(b *testing.B) { + for _, id := range tls13CipherSuites { + b.Run(tls.CipherSuiteName(id), func(b *testing.B) { benchmarkHKDFExpandLabel(b, id, true) }) + } +} + +func BenchmarkHKDFExpandLabelOurs(b *testing.B) { + for _, id := range tls13CipherSuites { + b.Run(tls.CipherSuiteName(id), func(b *testing.B) { benchmarkHKDFExpandLabel(b, id, false) }) + } +} + +func benchmarkHKDFExpandLabel(b *testing.B, cipherSuite uint16, useStdLib bool) { + b.ReportAllocs() + cs := cipherSuiteTLS13ByID(cipherSuite) + secret := make([]byte, 32) + rand.Read(secret) + + for b.Loop() { + if useStdLib { + nextTrafficSecret(cs, secret) + } else { + hkdfExpandLabel(cs.Hash, secret, nil, "traffic upd", cs.Hash.Size()) + } + } +} diff --git a/third_party/quic-go/internal/handshake/initial_aead.go b/third_party/quic-go/internal/handshake/initial_aead.go new file mode 100644 index 0000000..4b674d4 --- /dev/null +++ b/third_party/quic-go/internal/handshake/initial_aead.go @@ -0,0 +1,80 @@ +package handshake + +import ( + "crypto" + "crypto/hkdf" + "crypto/tls" + "fmt" + + "github.com/apernet/quic-go/internal/protocol" +) + +var ( + quicSaltV1 = []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a} + quicSaltV2 = []byte{0x0d, 0xed, 0xe3, 0xde, 0xf7, 0x00, 0xa6, 0xdb, 0x81, 0x93, 0x81, 0xbe, 0x6e, 0x26, 0x9d, 0xcb, 0xf9, 0xbd, 0x2e, 0xd9} +) + +const ( + hkdfLabelKeyV1 = "quic key" + hkdfLabelKeyV2 = "quicv2 key" + hkdfLabelIVV1 = "quic iv" + hkdfLabelIVV2 = "quicv2 iv" +) + +var initialSuite = getCipherSuite(tls.TLS_AES_128_GCM_SHA256) + +// NewInitialAEAD creates a new AEAD for Initial encryption / decryption. +func NewInitialAEAD(connID protocol.ConnectionID, pers protocol.Perspective, v protocol.Version) (LongHeaderSealer, LongHeaderOpener) { + var sealer LongHeaderSealer + var opener LongHeaderOpener + // The keys for the Initial AEAD are derived from the connection ID and constants defined in RFC 9001, Section 5.2. + // By design, the Initial encryption level provides no confidentiality against any attacker who has read the RFC. + // Its sole purpose is integrity protection. The Initial encryption level is therefore out of scope for FIPS 140. + // See also this thread on the IETF QUIC mailing list: https://mailarchive.ietf.org/arch/msg/quic/k2kl2W_n5WDEZBbt3O31Ef2XBbM. + withoutFIPSEnforcement(func() { + clientSecret, serverSecret := computeSecrets(connID, v) + var mySecret, otherSecret []byte + if pers == protocol.PerspectiveClient { + mySecret = clientSecret + otherSecret = serverSecret + } else { + mySecret = serverSecret + otherSecret = clientSecret + } + myKey, myIV := computeInitialKeyAndIV(mySecret, v) + otherKey, otherIV := computeInitialKeyAndIV(otherSecret, v) + + encrypter := initialSuite.AEAD(myKey, myIV) + decrypter := initialSuite.AEAD(otherKey, otherIV) + + sealer = newLongHeaderSealer(encrypter, newHeaderProtector(initialSuite, mySecret, true, v)) + opener = newLongHeaderOpener(decrypter, newHeaderProtector(initialSuite, otherSecret, true, v)) + }) + return sealer, opener +} + +func computeSecrets(connID protocol.ConnectionID, v protocol.Version) (clientSecret, serverSecret []byte) { + salt := quicSaltV1 + if v == protocol.Version2 { + salt = quicSaltV2 + } + initialSecret, err := hkdf.Extract(crypto.SHA256.New, connID.Bytes(), salt) + if err != nil { + panic(fmt.Errorf("quic: HKDF-Extract invocation failed unexpectedly: %v", err)) + } + clientSecret = hkdfExpandLabel(crypto.SHA256, initialSecret, []byte{}, "client in", crypto.SHA256.Size()) + serverSecret = hkdfExpandLabel(crypto.SHA256, initialSecret, []byte{}, "server in", crypto.SHA256.Size()) + return +} + +func computeInitialKeyAndIV(secret []byte, v protocol.Version) (key, iv []byte) { + keyLabel := hkdfLabelKeyV1 + ivLabel := hkdfLabelIVV1 + if v == protocol.Version2 { + keyLabel = hkdfLabelKeyV2 + ivLabel = hkdfLabelIVV2 + } + key = hkdfExpandLabel(crypto.SHA256, secret, []byte{}, keyLabel, 16) + iv = hkdfExpandLabel(crypto.SHA256, secret, []byte{}, ivLabel, 12) + return +} diff --git a/third_party/quic-go/internal/handshake/initial_aead_test.go b/third_party/quic-go/internal/handshake/initial_aead_test.go new file mode 100644 index 0000000..d83a0ca --- /dev/null +++ b/third_party/quic-go/internal/handshake/initial_aead_test.go @@ -0,0 +1,317 @@ +package handshake + +import ( + "bytes" + "crypto/rand" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestComputeClientKeyAndIV(t *testing.T) { + connID := protocol.ParseConnectionID(splitHexString(t, "0x8394c8f03e515708")) + + tests := []struct { + name string + version protocol.Version + expectedClientSecret []byte + expectedKey []byte + expectedIV []byte + }{ + { + name: "QUIC v1", + version: protocol.Version1, + expectedClientSecret: splitHexString(t, "c00cf151ca5be075ed0ebfb5c80323c4 2d6b7db67881289af4008f1f6c357aea"), + expectedKey: splitHexString(t, "1f369613dd76d5467730efcbe3b1a22d"), + expectedIV: splitHexString(t, "fa044b2f42a3fd3b46fb255c"), + }, + { + name: "QUIC v2", + version: protocol.Version2, + expectedClientSecret: splitHexString(t, "14ec9d6eb9fd7af83bf5a668bc17a7e2 83766aade7ecd0891f70f9ff7f4bf47b"), + expectedKey: splitHexString(t, "8b1a0bc121284290a29e0971b5cd045d"), + expectedIV: splitHexString(t, "91f73e2351d8fa91660e909f"), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + clientSecret, _ := computeSecrets(connID, tt.version) + require.Equal(t, tt.expectedClientSecret, clientSecret) + key, iv := computeInitialKeyAndIV(clientSecret, tt.version) + require.Equal(t, tt.expectedKey, key) + require.Equal(t, tt.expectedIV, iv) + }) + } +} + +func TestComputeServerKeyAndIV(t *testing.T) { + connID := protocol.ParseConnectionID(splitHexString(t, "0x8394c8f03e515708")) + + tests := []struct { + name string + version protocol.Version + expectedServerSecret []byte + expectedKey []byte + expectedIV []byte + }{ + { + name: "QUIC v1", + version: protocol.Version1, + expectedServerSecret: splitHexString(t, "3c199828fd139efd216c155ad844cc81 fb82fa8d7446fa7d78be803acdda951b"), + expectedKey: splitHexString(t, "cf3a5331653c364c88f0f379b6067e37"), + expectedIV: splitHexString(t, "0ac1493ca1905853b0bba03e"), + }, + { + name: "QUIC v2", + version: protocol.Version2, + expectedServerSecret: splitHexString(t, "0263db1782731bf4588e7e4d93b74639 07cb8cd8200b5da55a8bd488eafc37c1"), + expectedKey: splitHexString(t, "82db637861d55e1d011f19ea71d5d2a7"), + expectedIV: splitHexString(t, "dd13c276499c0249d3310652"), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, serverSecret := computeSecrets(connID, tt.version) + require.Equal(t, tt.expectedServerSecret, serverSecret) + key, iv := computeInitialKeyAndIV(serverSecret, tt.version) + require.Equal(t, tt.expectedKey, key) + require.Equal(t, tt.expectedIV, iv) + }) + } +} + +func TestClientInitial(t *testing.T) { + connID := protocol.ParseConnectionID(splitHexString(t, "0x8394c8f03e515708")) + + tests := []struct { + version protocol.Version + header []byte + data []byte + expectedSample []byte + expectedHdrFirstByte byte + expectedHdr []byte + expectedPacket []byte + }{ + { + version: protocol.Version1, + header: splitHexString(t, "c300000001088394c8f03e5157080000449e00000002"), + data: splitHexString(t, "060040f1010000ed0303ebf8fa56f129 39b9584a3896472ec40bb863cfd3e868 04fe3a47f06a2b69484c000004130113 02010000c000000010000e00000b6578 616d706c652e636f6dff01000100000a 00080006001d00170018001000070005 04616c706e0005000501000000000033 00260024001d00209370b2c9caa47fba baf4559fedba753de171fa71f50f1ce1 5d43e994ec74d748002b000302030400 0d0010000e0403050306030203080408 050806002d00020101001c0002400100 3900320408ffffffffffffffff050480 00ffff07048000ffff08011001048000 75300901100f088394c8f03e51570806 048000ffff"), + expectedSample: splitHexString(t, "d1b1c98dd7689fb8ec11d242b123dc9b"), + expectedHdrFirstByte: byte(0xc0), + expectedHdr: splitHexString(t, "7b9aec34"), + expectedPacket: splitHexString(t, "c000000001088394c8f03e5157080000 449e7b9aec34d1b1c98dd7689fb8ec11 d242b123dc9bd8bab936b47d92ec356c 0bab7df5976d27cd449f63300099f399 1c260ec4c60d17b31f8429157bb35a12 82a643a8d2262cad67500cadb8e7378c 8eb7539ec4d4905fed1bee1fc8aafba1 7c750e2c7ace01e6005f80fcb7df6212 30c83711b39343fa028cea7f7fb5ff89 eac2308249a02252155e2347b63d58c5 457afd84d05dfffdb20392844ae81215 4682e9cf012f9021a6f0be17ddd0c208 4dce25ff9b06cde535d0f920a2db1bf3 62c23e596d11a4f5a6cf3948838a3aec 4e15daf8500a6ef69ec4e3feb6b1d98e 610ac8b7ec3faf6ad760b7bad1db4ba3 485e8a94dc250ae3fdb41ed15fb6a8e5 eba0fc3dd60bc8e30c5c4287e53805db 059ae0648db2f64264ed5e39be2e20d8 2df566da8dd5998ccabdae053060ae6c 7b4378e846d29f37ed7b4ea9ec5d82e7 961b7f25a9323851f681d582363aa5f8 9937f5a67258bf63ad6f1a0b1d96dbd4 faddfcefc5266ba6611722395c906556 be52afe3f565636ad1b17d508b73d874 3eeb524be22b3dcbc2c7468d54119c74 68449a13d8e3b95811a198f3491de3e7 fe942b330407abf82a4ed7c1b311663a c69890f4157015853d91e923037c227a 33cdd5ec281ca3f79c44546b9d90ca00 f064c99e3dd97911d39fe9c5d0b23a22 9a234cb36186c4819e8b9c5927726632 291d6a418211cc2962e20fe47feb3edf 330f2c603a9d48c0fcb5699dbfe58964 25c5bac4aee82e57a85aaf4e2513e4f0 5796b07ba2ee47d80506f8d2c25e50fd 14de71e6c418559302f939b0e1abd576 f279c4b2e0feb85c1f28ff18f58891ff ef132eef2fa09346aee33c28eb130ff2 8f5b766953334113211996d20011a198 e3fc433f9f2541010ae17c1bf202580f 6047472fb36857fe843b19f5984009dd c324044e847a4f4a0ab34f719595de37 252d6235365e9b84392b061085349d73 203a4a13e96f5432ec0fd4a1ee65accd d5e3904df54c1da510b0ff20dcc0c77f cb2c0e0eb605cb0504db87632cf3d8b4 dae6e705769d1de354270123cb11450e fc60ac47683d7b8d0f811365565fd98c 4c8eb936bcab8d069fc33bd801b03ade a2e1fbc5aa463d08ca19896d2bf59a07 1b851e6c239052172f296bfb5e724047 90a2181014f3b94a4e97d117b4381303 68cc39dbb2d198065ae3986547926cd2 162f40a29f0c3c8745c0f50fba3852e5 66d44575c29d39a03f0cda721984b6f4 40591f355e12d439ff150aab7613499d bd49adabc8676eef023b15b65bfc5ca0 6948109f23f350db82123535eb8a7433 bdabcb909271a6ecbcb58b936a88cd4e 8f2e6ff5800175f113253d8fa9ca8885 c2f552e657dc603f252e1a8e308f76f0 be79e2fb8f5d5fbbe2e30ecadd220723 c8c0aea8078cdfcb3868263ff8f09400 54da48781893a7e49ad5aff4af300cd8 04a6b6279ab3ff3afb64491c85194aab 760d58a606654f9f4400e8b38591356f bf6425aca26dc85244259ff2b19c41b9 f96f3ca9ec1dde434da7d2d392b905dd f3d1f9af93d1af5950bd493f5aa731b4 056df31bd267b6b90a079831aaf579be 0a39013137aac6d404f518cfd4684064 7e78bfe706ca4cf5e9c5453e9f7cfd2b 8b4c8d169a44e55c88d4a9a7f9474241 e221af44860018ab0856972e194cd934"), + }, + { + version: protocol.Version2, + header: splitHexString(t, "d36b3343cf088394c8f03e5157080000449e00000002"), + data: splitHexString(t, "060040f1010000ed0303ebf8fa56f129 39b9584a3896472ec40bb863cfd3e868 04fe3a47f06a2b69484c000004130113 02010000c000000010000e00000b6578 616d706c652e636f6dff01000100000a 00080006001d00170018001000070005 04616c706e0005000501000000000033 00260024001d00209370b2c9caa47fba baf4559fedba753de171fa71f50f1ce1 5d43e994ec74d748002b000302030400 0d0010000e0403050306030203080408 050806002d00020101001c0002400100 3900320408ffffffffffffffff050480 00ffff07048000ffff08011001048000 75300901100f088394c8f03e51570806 048000ffff"), + expectedSample: splitHexString(t, "ffe67b6abcdb4298b485dd04de806071"), + expectedHdrFirstByte: byte(0xd7), + expectedHdr: splitHexString(t, "a0c95e82"), + expectedPacket: splitHexString(t, "d76b3343cf088394c8f03e5157080000 449ea0c95e82ffe67b6abcdb4298b485 dd04de806071bf03dceebfa162e75d6c 96058bdbfb127cdfcbf903388e99ad04 9f9a3dd4425ae4d0992cfff18ecf0fdb 5a842d09747052f17ac2053d21f57c5d 250f2c4f0e0202b70785b7946e992e58 a59ac52dea6774d4f03b55545243cf1a 12834e3f249a78d395e0d18f4d766004 f1a2674802a747eaa901c3f10cda5500 cb9122faa9f1df66c392079a1b40f0de 1c6054196a11cbea40afb6ef5253cd68 18f6625efce3b6def6ba7e4b37a40f77 32e093daa7d52190935b8da58976ff33 12ae50b187c1433c0f028edcc4c2838b 6a9bfc226ca4b4530e7a4ccee1bfa2a3 d396ae5a3fb512384b2fdd851f784a65 e03f2c4fbe11a53c7777c023462239dd 6f7521a3f6c7d5dd3ec9b3f233773d4b 46d23cc375eb198c63301c21801f6520 bcfb7966fc49b393f0061d974a2706df 8c4a9449f11d7f3d2dcbb90c6b877045 636e7c0c0fe4eb0f697545460c806910 d2c355f1d253bc9d2452aaa549e27a1f ac7cf4ed77f322e8fa894b6a83810a34 b361901751a6f5eb65a0326e07de7c12 16ccce2d0193f958bb3850a833f7ae43 2b65bc5a53975c155aa4bcb4f7b2c4e5 4df16efaf6ddea94e2c50b4cd1dfe060 17e0e9d02900cffe1935e0491d77ffb4 fdf85290fdd893d577b1131a610ef6a5 c32b2ee0293617a37cbb08b847741c3b 8017c25ca9052ca1079d8b78aebd4787 6d330a30f6a8c6d61dd1ab5589329de7 14d19d61370f8149748c72f132f0fc99 f34d766c6938597040d8f9e2bb522ff9 9c63a344d6a2ae8aa8e51b7b90a4a806 105fcbca31506c446151adfeceb51b91 abfe43960977c87471cf9ad4074d30e1 0d6a7f03c63bd5d4317f68ff325ba3bd 80bf4dc8b52a0ba031758022eb025cdd 770b44d6d6cf0670f4e990b22347a7db 848265e3e5eb72dfe8299ad7481a4083 22cac55786e52f633b2fb6b614eaed18 d703dd84045a274ae8bfa73379661388 d6991fe39b0d93debb41700b41f90a15 c4d526250235ddcd6776fc77bc97e7a4 17ebcb31600d01e57f32162a8560cacc 7e27a096d37a1a86952ec71bd89a3e9a 30a2a26162984d7740f81193e8238e61 f6b5b984d4d3dfa033c1bb7e4f0037fe bf406d91c0dccf32acf423cfa1e70710 10d3f270121b493ce85054ef58bada42 310138fe081adb04e2bd901f2f13458b 3d6758158197107c14ebb193230cd115 7380aa79cae1374a7c1e5bbcb80ee23e 06ebfde206bfb0fcbc0edc4ebec30966 1bdd908d532eb0c6adc38b7ca7331dce 8dfce39ab71e7c32d318d136b6100671 a1ae6a6600e3899f31f0eed19e3417d1 34b90c9058f8632c798d4490da498730 7cba922d61c39805d072b589bd52fdf1 e86215c2d54e6670e07383a27bbffb5a ddf47d66aa85a0c6f9f32e59d85a44dd 5d3b22dc2be80919b490437ae4f36a0a e55edf1d0b5cb4e9a3ecabee93dfc6e3 8d209d0fa6536d27a5d6fbb17641cde2 7525d61093f1b28072d111b2b4ae5f89 d5974ee12e5cf7d5da4d6a31123041f3 3e61407e76cffcdcfd7e19ba58cf4b53 6f4c4938ae79324dc402894b44faf8af bab35282ab659d13c93f70412e85cb19 9a37ddec600545473cfb5a05e08d0b20 9973b2172b4d21fb69745a262ccde96b a18b2faa745b6fe189cf772a9f84cbfc"), + }, + } + + for _, tt := range tests { + t.Run(tt.version.String(), func(t *testing.T) { + sealer, _ := NewInitialAEAD(connID, protocol.PerspectiveClient, tt.version) + tt.data = append(tt.data, make([]byte, 1162-len(tt.data))...) // add PADDING + sealed := sealer.Seal(nil, tt.data, 2, tt.header) + sample := sealed[0:16] + require.Equal(t, tt.expectedSample, sample) + sealer.EncryptHeader(sample, &tt.header[0], tt.header[len(tt.header)-4:]) + require.Equal(t, tt.expectedHdrFirstByte, tt.header[0]) + require.Equal(t, tt.expectedHdr, tt.header[len(tt.header)-4:]) + packet := append(tt.header, sealed...) + require.Equal(t, tt.expectedPacket, packet) + }) + } +} + +func TestServersInitial(t *testing.T) { + connID := protocol.ParseConnectionID(splitHexString(t, "0x8394c8f03e515708")) + + testCases := []struct { + name string + version protocol.Version + header []byte + data []byte + expectedSample []byte + expectedHdr []byte + expectedPacket []byte + }{ + { + name: "QUIC v1", + version: protocol.Version1, + header: splitHexString(t, "c1000000010008f067a5502a4262b50040750001"), + data: splitHexString(t, "02000000000600405a020000560303ee fce7f7b37ba1d1632e96677825ddf739 88cfc79825df566dc5430b9a045a1200 130100002e00330024001d00209d3c94 0d89690b84d08a60993c144eca684d10 81287c834d5311bcf32bb9da1a002b00 020304"), + expectedSample: splitHexString(t, "2cd0991cd25b0aac406a5816b6394100"), + expectedHdr: splitHexString(t, "cf000000010008f067a5502a4262b5004075c0d9"), + expectedPacket: splitHexString(t, "cf000000010008f067a5502a4262b500 4075c0d95a482cd0991cd25b0aac406a 5816b6394100f37a1c69797554780bb3 8cc5a99f5ede4cf73c3ec2493a1839b3 dbcba3f6ea46c5b7684df3548e7ddeb9 c3bf9c73cc3f3bded74b562bfb19fb84 022f8ef4cdd93795d77d06edbb7aaf2f 58891850abbdca3d20398c276456cbc4 2158407dd074ee"), + }, + { + name: "QUIC v2", + version: protocol.Version2, + header: splitHexString(t, "d16b3343cf0008f067a5502a4262b50040750001"), + data: splitHexString(t, "02000000000600405a020000560303ee fce7f7b37ba1d1632e96677825ddf739 88cfc79825df566dc5430b9a045a1200 130100002e00330024001d00209d3c94 0d89690b84d08a60993c144eca684d10 81287c834d5311bcf32bb9da1a002b00 020304"), + expectedSample: splitHexString(t, "6f05d8a4398c47089698baeea26b91eb"), + expectedHdr: splitHexString(t, "dc6b3343cf0008f067a5502a4262b5004075d92f"), + expectedPacket: splitHexString(t, "dc6b3343cf0008f067a5502a4262b500 4075d92faaf16f05d8a4398c47089698 baeea26b91eb761d9b89237bbf872630 17915358230035f7fd3945d88965cf17 f9af6e16886c61bfc703106fbaf3cb4c fa52382dd16a393e42757507698075b2 c984c707f0a0812d8cd5a6881eaf21ce da98f4bd23f6fe1a3e2c43edd9ce7ca8 4bed8521e2e140"), + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + sealer, _ := NewInitialAEAD(connID, protocol.PerspectiveServer, tc.version) + sealed := sealer.Seal(nil, tc.data, 1, tc.header) + sample := sealed[2 : 2+16] + require.Equal(t, tc.expectedSample, sample) + sealer.EncryptHeader(sample, &tc.header[0], tc.header[len(tc.header)-2:]) + require.Equal(t, tc.expectedHdr, tc.header) + packet := append(tc.header, sealed...) + require.Equal(t, tc.expectedPacket, packet) + }) + } +} + +func TestInitialAEADSealsAndOpens(t *testing.T) { + for _, ver := range []protocol.Version{protocol.Version1, protocol.Version2} { + t.Run(ver.String(), func(t *testing.T) { + connectionID := protocol.ParseConnectionID([]byte{0x12, 0x34, 0x56, 0x78, 0x90, 0xab, 0xcd, 0xef}) + clientSealer, clientOpener := NewInitialAEAD(connectionID, protocol.PerspectiveClient, ver) + serverSealer, serverOpener := NewInitialAEAD(connectionID, protocol.PerspectiveServer, ver) + + clientMessage := clientSealer.Seal(nil, []byte("foobar"), 42, []byte("aad")) + m, err := serverOpener.Open(nil, clientMessage, 42, []byte("aad")) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), m) + serverMessage := serverSealer.Seal(nil, []byte("raboof"), 99, []byte("daa")) + m, err = clientOpener.Open(nil, serverMessage, 99, []byte("daa")) + require.NoError(t, err) + require.Equal(t, []byte("raboof"), m) + }) + } +} + +func TestInitialAEADFailsWithDifferentConnectionIDs(t *testing.T) { + for _, ver := range []protocol.Version{protocol.Version1, protocol.Version2} { + t.Run(ver.String(), func(t *testing.T) { + c1 := protocol.ParseConnectionID([]byte{0, 0, 0, 0, 0, 0, 0, 1}) + c2 := protocol.ParseConnectionID([]byte{0, 0, 0, 0, 0, 0, 0, 2}) + clientSealer, _ := NewInitialAEAD(c1, protocol.PerspectiveClient, ver) + _, serverOpener := NewInitialAEAD(c2, protocol.PerspectiveServer, ver) + + clientMessage := clientSealer.Seal(nil, []byte("foobar"), 42, []byte("aad")) + _, err := serverOpener.Open(nil, clientMessage, 42, []byte("aad")) + require.Equal(t, ErrDecryptionFailed, err) + }) + } +} + +func TestInitialAEADEncryptsAndDecryptsHeader(t *testing.T) { + for _, ver := range []protocol.Version{protocol.Version1, protocol.Version2} { + t.Run(ver.String(), func(t *testing.T) { + connID := protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}) + clientSealer, clientOpener := NewInitialAEAD(connID, protocol.PerspectiveClient, ver) + serverSealer, serverOpener := NewInitialAEAD(connID, protocol.PerspectiveServer, ver) + + header := []byte{0x5e, 0, 1, 2, 3, 4, 0xde, 0xad, 0xbe, 0xef} + sample := make([]byte, 16) + rand.Read(sample) + clientSealer.EncryptHeader(sample, &header[0], header[6:10]) + require.Equal(t, byte(0x5e&0xf0), header[0]&0xf0) + require.Equal(t, []byte{0, 1, 2, 3, 4}, header[1:6]) + require.NotEqual(t, []byte{0xde, 0xad, 0xbe, 0xef}, header[6:10]) + serverOpener.DecryptHeader(sample, &header[0], header[6:10]) + require.Equal(t, byte(0x5e), header[0]) + require.Equal(t, []byte{0, 1, 2, 3, 4}, header[1:6]) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, header[6:10]) + + serverSealer.EncryptHeader(sample, &header[0], header[6:10]) + require.Equal(t, byte(0x5e&0xf0), header[0]&0xf0) + require.Equal(t, []byte{0, 1, 2, 3, 4}, header[1:6]) + require.NotEqual(t, []byte{0xde, 0xad, 0xbe, 0xef}, header[6:10]) + clientOpener.DecryptHeader(sample, &header[0], header[6:10]) + require.Equal(t, byte(0x5e), header[0]) + require.Equal(t, []byte{0, 1, 2, 3, 4}, header[1:6]) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, header[6:10]) + }) + } +} + +func BenchmarkInitialAEADCreate(b *testing.B) { + b.ReportAllocs() + connID := protocol.ParseConnectionID([]byte{0x12, 0x34, 0x56, 0x78, 0x90, 0xab, 0xcd, 0xef}) + + for b.Loop() { + NewInitialAEAD(connID, protocol.PerspectiveServer, protocol.Version1) + } +} + +func BenchmarkInitialAEAD(b *testing.B) { + connectionID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd}) + clientSealer, _ := NewInitialAEAD(connectionID, protocol.PerspectiveClient, protocol.Version1) + _, serverOpener := NewInitialAEAD(connectionID, protocol.PerspectiveServer, protocol.Version1) + + packetData := make([]byte, 1200) + rand.Read(packetData) + hdr := make([]byte, 50) + rand.Read(hdr) + msg := clientSealer.Seal(nil, packetData, 42, hdr) + m, err := serverOpener.Open(nil, msg, 42, hdr) + if err != nil { + b.Fatalf("opening failed: %s", err) + } + if !bytes.Equal(m, packetData) { + b.Fatal("decrypted data doesn't match") + } + + b.ResetTimer() + b.Run("opening 100 bytes", func(b *testing.B) { + sealer, _ := NewInitialAEAD(connectionID, protocol.PerspectiveClient, protocol.Version1) + _, opener := NewInitialAEAD(connectionID, protocol.PerspectiveServer, protocol.Version1) + benchmarkOpen(b, opener, sealer.Seal(nil, packetData[:100], 42, hdr), 42, hdr) + }) + b.Run("opening 1200 bytes", func(b *testing.B) { + sealer, _ := NewInitialAEAD(connectionID, protocol.PerspectiveClient, protocol.Version1) + _, opener := NewInitialAEAD(connectionID, protocol.PerspectiveServer, protocol.Version1) + benchmarkOpen(b, opener, sealer.Seal(nil, packetData, 42, hdr), 42, hdr) + }) + + b.Run("sealing 100 bytes", func(b *testing.B) { + sealer, _ := NewInitialAEAD(connectionID, protocol.PerspectiveClient, protocol.Version1) + benchmarkSeal(b, sealer, packetData[:100], hdr) + }) + b.Run("sealing 1200 bytes", func(b *testing.B) { + sealer, _ := NewInitialAEAD(connectionID, protocol.PerspectiveClient, protocol.Version1) + benchmarkSeal(b, sealer, packetData, hdr) + }) +} + +func benchmarkOpen(b *testing.B, aead LongHeaderOpener, msg []byte, pn protocol.PacketNumber, hdr []byte) { + b.ReportAllocs() + dst := make([]byte, 0, 1500) + + for b.Loop() { + dst = dst[:0] + if _, err := aead.Open(dst, msg, pn, hdr); err != nil { + b.Fatalf("opening failed: %s", err) + } + } +} + +func benchmarkSeal(b *testing.B, aead LongHeaderSealer, msg, hdr []byte) { + b.ReportAllocs() + dst := make([]byte, 0, 1500) + + var pn protocol.PacketNumber + for b.Loop() { + dst = dst[:0] + aead.Seal(dst, msg, pn, hdr) + pn++ + } +} diff --git a/third_party/quic-go/internal/handshake/interface.go b/third_party/quic-go/internal/handshake/interface.go new file mode 100644 index 0000000..7dd1228 --- /dev/null +++ b/third_party/quic-go/internal/handshake/interface.go @@ -0,0 +1,140 @@ +package handshake + +import ( + "context" + "crypto/tls" + "errors" + "io" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +var ( + // ErrKeysNotYetAvailable is returned when an opener or a sealer is requested for an encryption level, + // but the corresponding opener has not yet been initialized + // This can happen when packets arrive out of order. + ErrKeysNotYetAvailable = errors.New("CryptoSetup: keys at this encryption level not yet available") + // ErrKeysDropped is returned when an opener or a sealer is requested for an encryption level, + // but the corresponding keys have already been dropped. + ErrKeysDropped = errors.New("CryptoSetup: keys were already dropped") + // ErrDecryptionFailed is returned when the AEAD fails to open the packet. + ErrDecryptionFailed = errors.New("decryption failed") +) + +type headerDecryptor interface { + DecryptHeader(sample []byte, firstByte *byte, pnBytes []byte) +} + +// LongHeaderOpener opens a long header packet +type LongHeaderOpener interface { + headerDecryptor + DecodePacketNumber(wirePN protocol.PacketNumber, wirePNLen protocol.PacketNumberLen) protocol.PacketNumber + Open(dst, src []byte, pn protocol.PacketNumber, associatedData []byte) ([]byte, error) +} + +// ShortHeaderOpener opens a short header packet +type ShortHeaderOpener interface { + headerDecryptor + DecodePacketNumber(wirePN protocol.PacketNumber, wirePNLen protocol.PacketNumberLen) protocol.PacketNumber + Open(dst, src []byte, rcvTime monotime.Time, pn protocol.PacketNumber, kp protocol.KeyPhaseBit, associatedData []byte) ([]byte, error) +} + +// LongHeaderSealer seals a long header packet +type LongHeaderSealer interface { + Seal(dst, src []byte, packetNumber protocol.PacketNumber, associatedData []byte) []byte + EncryptHeader(sample []byte, firstByte *byte, pnBytes []byte) + Overhead() int +} + +// ShortHeaderSealer seals a short header packet +type ShortHeaderSealer interface { + LongHeaderSealer + KeyPhase() protocol.KeyPhaseBit +} + +type ConnectionState struct { + tls.ConnectionState + Used0RTT bool +} + +// EventKind is the kind of handshake event. +type EventKind uint8 + +const ( + // EventNoEvent signals that there are no new handshake events + EventNoEvent EventKind = iota + 1 + // EventWriteInitialData contains new CRYPTO data to send at the Initial encryption level + EventWriteInitialData + // EventWriteHandshakeData contains new CRYPTO data to send at the Handshake encryption level + EventWriteHandshakeData + // EventReceivedReadKeys signals that new decryption keys are available. + // It doesn't say which encryption level those keys are for. + EventReceivedReadKeys + // EventDiscard0RTTKeys signals that the Handshake keys were discarded. + EventDiscard0RTTKeys + // EventReceivedTransportParameters contains the transport parameters sent by the peer. + EventReceivedTransportParameters + // EventRestoredTransportParameters contains the transport parameters restored from the session ticket. + // It is only used for the client. + EventRestoredTransportParameters + // EventHandshakeComplete signals that the TLS handshake was completed. + EventHandshakeComplete +) + +func (k EventKind) String() string { + switch k { + case EventNoEvent: + return "EventNoEvent" + case EventWriteInitialData: + return "EventWriteInitialData" + case EventWriteHandshakeData: + return "EventWriteHandshakeData" + case EventReceivedReadKeys: + return "EventReceivedReadKeys" + case EventDiscard0RTTKeys: + return "EventDiscard0RTTKeys" + case EventReceivedTransportParameters: + return "EventReceivedTransportParameters" + case EventRestoredTransportParameters: + return "EventRestoredTransportParameters" + case EventHandshakeComplete: + return "EventHandshakeComplete" + default: + return "Unknown EventKind" + } +} + +// Event is a handshake event. +type Event struct { + Kind EventKind + Data []byte + TransportParameters *wire.TransportParameters +} + +// CryptoSetup handles the handshake and protecting / unprotecting packets +type CryptoSetup interface { + StartHandshake(context.Context) error + io.Closer + ChangeConnectionID(protocol.ConnectionID) + GetSessionTicket() ([]byte, error) + + HandleMessage([]byte, protocol.EncryptionLevel) error + NextEvent() Event + + SetLargest1RTTAcked(protocol.PacketNumber) error + DiscardInitialKeys() + SetHandshakeConfirmed() + ConnectionState() ConnectionState + + GetInitialOpener() (LongHeaderOpener, error) + GetHandshakeOpener() (LongHeaderOpener, error) + Get0RTTOpener() (LongHeaderOpener, error) + Get1RTTOpener() (ShortHeaderOpener, error) + + GetInitialSealer() (LongHeaderSealer, error) + GetHandshakeSealer() (LongHeaderSealer, error) + Get0RTTSealer() (LongHeaderSealer, error) + Get1RTTSealer() (ShortHeaderSealer, error) +} diff --git a/third_party/quic-go/internal/handshake/quic_event_go125.go b/third_party/quic-go/internal/handshake/quic_event_go125.go new file mode 100644 index 0000000..e955397 --- /dev/null +++ b/third_party/quic-go/internal/handshake/quic_event_go125.go @@ -0,0 +1,11 @@ +//go:build go1.25 && !go1.26 + +package handshake + +import "crypto/tls" + +const quicErrorEvent tls.QUICEventKind = -1 + +func extractQUICEventError(tls.QUICEvent) error { + return nil +} diff --git a/third_party/quic-go/internal/handshake/quic_event_go126.go b/third_party/quic-go/internal/handshake/quic_event_go126.go new file mode 100644 index 0000000..734778e --- /dev/null +++ b/third_party/quic-go/internal/handshake/quic_event_go126.go @@ -0,0 +1,11 @@ +//go:build go1.26 + +package handshake + +import "crypto/tls" + +const quicErrorEvent tls.QUICEventKind = tls.QUICErrorEvent + +func extractQUICEventError(ev tls.QUICEvent) error { + return ev.Err +} diff --git a/third_party/quic-go/internal/handshake/retry_go125.go b/third_party/quic-go/internal/handshake/retry_go125.go new file mode 100644 index 0000000..c8a9454 --- /dev/null +++ b/third_party/quic-go/internal/handshake/retry_go125.go @@ -0,0 +1,68 @@ +//go:build !go1.26 + +package handshake + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "fmt" + "sync" + + "github.com/apernet/quic-go/internal/protocol" +) + +// Instead of using an init function, the AEADs are created lazily. +// For more details see https://github.com/apernet/quic-go/issues/4894. +var ( + retryAEADv1 cipher.AEAD // used for QUIC v1 (RFC 9000) + retryAEADv2 cipher.AEAD // used for QUIC v2 (RFC 9369) +) + +func initAEAD(key [16]byte) cipher.AEAD { + aes, err := aes.NewCipher(key[:]) + if err != nil { + panic(err) + } + aead, err := cipher.NewGCM(aes) + if err != nil { + panic(err) + } + return aead +} + +var ( + retryBuf bytes.Buffer + retryMutex sync.Mutex + retryNonceV1 = [12]byte{0x46, 0x15, 0x99, 0xd3, 0x5d, 0x63, 0x2b, 0xf2, 0x23, 0x98, 0x25, 0xbb} + retryNonceV2 = [12]byte{0xd8, 0x69, 0x69, 0xbc, 0x2d, 0x7c, 0x6d, 0x99, 0x90, 0xef, 0xb0, 0x4a} +) + +// GetRetryIntegrityTag calculates the integrity tag on a Retry packet +func GetRetryIntegrityTag(retry []byte, origDestConnID protocol.ConnectionID, version protocol.Version) *[16]byte { + retryMutex.Lock() + defer retryMutex.Unlock() + + retryBuf.WriteByte(uint8(origDestConnID.Len())) + retryBuf.Write(origDestConnID.Bytes()) + retryBuf.Write(retry) + defer retryBuf.Reset() + + var tag [16]byte + var sealed []byte + if version == protocol.Version2 { + if retryAEADv2 == nil { + retryAEADv2 = initAEAD([16]byte{0x8f, 0xb4, 0xb0, 0x1b, 0x56, 0xac, 0x48, 0xe2, 0x60, 0xfb, 0xcb, 0xce, 0xad, 0x7c, 0xcc, 0x92}) + } + sealed = retryAEADv2.Seal(tag[:0], retryNonceV2[:], nil, retryBuf.Bytes()) + } else { + if retryAEADv1 == nil { + retryAEADv1 = initAEAD([16]byte{0xbe, 0x0c, 0x69, 0x0b, 0x9f, 0x66, 0x57, 0x5a, 0x1d, 0x76, 0x6b, 0x54, 0xe3, 0x68, 0xc8, 0x4e}) + } + sealed = retryAEADv1.Seal(tag[:0], retryNonceV1[:], nil, retryBuf.Bytes()) + } + if len(sealed) != 16 { + panic(fmt.Sprintf("unexpected Retry integrity tag length: %d", len(sealed))) + } + return &tag +} diff --git a/third_party/quic-go/internal/handshake/retry_go126.go b/third_party/quic-go/internal/handshake/retry_go126.go new file mode 100644 index 0000000..57004c8 --- /dev/null +++ b/third_party/quic-go/internal/handshake/retry_go126.go @@ -0,0 +1,70 @@ +//go:build go1.26 + +package handshake + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "crypto/fips140" + "fmt" + "sync" + + "github.com/apernet/quic-go/internal/protocol" +) + +// used for QUIC v1 (RFC 9000) +var retryAEADv1 = initAEAD([16]byte{0xbe, 0x0c, 0x69, 0x0b, 0x9f, 0x66, 0x57, 0x5a, 0x1d, 0x76, 0x6b, 0x54, 0xe3, 0x68, 0xc8, 0x4e}) + +// used for QUIC v2 (RFC 9369) +var retryAEADv2 = initAEAD([16]byte{0x8f, 0xb4, 0xb0, 0x1b, 0x56, 0xac, 0x48, 0xe2, 0x60, 0xfb, 0xcb, 0xce, 0xad, 0x7c, 0xcc, 0x92}) + +func initAEAD(key [16]byte) cipher.AEAD { + aes, err := aes.NewCipher(key[:]) + if err != nil { + panic(err) + } + var aead cipher.AEAD + // Retry packet authentication uses the fixed key and nonce specified by RFC 9000. + // It acts as an integrity tag for the Retry packet itself, protecting against + // accidental modification and making injection harder. It is not used to encrypt + // packet contents and therefore outside the scope of FIPS 140 enforcement. + fips140.WithoutEnforcement(func() { + var err error + aead, err = cipher.NewGCM(aes) + if err != nil { + panic(err) + } + }) + return aead +} + +var ( + retryBuf bytes.Buffer + retryMutex sync.Mutex + retryNonceV1 = [12]byte{0x46, 0x15, 0x99, 0xd3, 0x5d, 0x63, 0x2b, 0xf2, 0x23, 0x98, 0x25, 0xbb} + retryNonceV2 = [12]byte{0xd8, 0x69, 0x69, 0xbc, 0x2d, 0x7c, 0x6d, 0x99, 0x90, 0xef, 0xb0, 0x4a} +) + +// GetRetryIntegrityTag calculates the integrity tag on a Retry packet +func GetRetryIntegrityTag(retry []byte, origDestConnID protocol.ConnectionID, version protocol.Version) *[16]byte { + retryMutex.Lock() + defer retryMutex.Unlock() + + retryBuf.WriteByte(uint8(origDestConnID.Len())) + retryBuf.Write(origDestConnID.Bytes()) + retryBuf.Write(retry) + defer retryBuf.Reset() + + var tag [16]byte + var sealed []byte + if version == protocol.Version2 { + sealed = retryAEADv2.Seal(tag[:0], retryNonceV2[:], nil, retryBuf.Bytes()) + } else { + sealed = retryAEADv1.Seal(tag[:0], retryNonceV1[:], nil, retryBuf.Bytes()) + } + if len(sealed) != 16 { + panic(fmt.Sprintf("unexpected Retry integrity tag length: %d", len(sealed))) + } + return &tag +} diff --git a/third_party/quic-go/internal/handshake/retry_test.go b/third_party/quic-go/internal/handshake/retry_test.go new file mode 100644 index 0000000..d2c8e87 --- /dev/null +++ b/third_party/quic-go/internal/handshake/retry_test.go @@ -0,0 +1,56 @@ +package handshake + +import ( + "encoding/binary" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestRetryIntegrityTagCalculation(t *testing.T) { + connID := protocol.ParseConnectionID([]byte{1, 2, 3, 4}) + fooTag := GetRetryIntegrityTag([]byte("foo"), connID, protocol.Version1) + barTag := GetRetryIntegrityTag([]byte("bar"), connID, protocol.Version1) + require.NotNil(t, fooTag) + require.NotNil(t, barTag) + require.NotEqual(t, *fooTag, *barTag) +} + +func TestRetryIntegrityTagWithDifferentConnectionIDs(t *testing.T) { + connID1 := protocol.ParseConnectionID([]byte{1, 2, 3, 4}) + connID2 := protocol.ParseConnectionID([]byte{4, 3, 2, 1}) + t1 := GetRetryIntegrityTag([]byte("foobar"), connID1, protocol.Version1) + t2 := GetRetryIntegrityTag([]byte("foobar"), connID2, protocol.Version1) + require.NotEqual(t, *t1, *t2) +} + +func TestRetryIntegrityTagWithTestVectors(t *testing.T) { + tests := []struct { + name string + version protocol.Version + data []byte + }{ + { + name: "v1", + version: protocol.Version1, + data: splitHexString(t, "ff000000010008f067a5502a4262b574 6f6b656e04a265ba2eff4d829058fb3f 0f2496ba"), + }, + { + name: "v2", + version: protocol.Version2, + data: splitHexString(t, "cf6b3343cf0008f067a5502a4262b574 6f6b656ec8646ce8bfe33952d9555436 65dcc7b6"), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + v := binary.BigEndian.Uint32(tt.data[1:5]) + require.Equal(t, tt.version, protocol.Version(v)) + connID := protocol.ParseConnectionID(splitHexString(t, "0x8394c8f03e515708")) + tag := GetRetryIntegrityTag(tt.data[:len(tt.data)-16], connID, tt.version) + require.Equal(t, tt.data[len(tt.data)-16:], tag[:]) + }) + } +} diff --git a/third_party/quic-go/internal/handshake/session_ticket.go b/third_party/quic-go/internal/handshake/session_ticket.go new file mode 100644 index 0000000..488921c --- /dev/null +++ b/third_party/quic-go/internal/handshake/session_ticket.go @@ -0,0 +1,55 @@ +package handshake + +import ( + "bytes" + "errors" + "fmt" + + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/quicvarint" +) + +const sessionTicketRevision = 5 + +type sessionTicket struct { + Parameters *wire.TransportParameters +} + +func (t *sessionTicket) Marshal() []byte { + b := make([]byte, 0, 256) + b = quicvarint.Append(b, sessionTicketRevision) + return t.Parameters.MarshalForSessionTicket(b) +} + +func (t *sessionTicket) Unmarshal(b []byte) error { + rev, l, err := quicvarint.Parse(b) + if err != nil { + return errors.New("failed to read session ticket revision") + } + b = b[l:] + if rev != sessionTicketRevision { + return fmt.Errorf("unknown session ticket revision: %d", rev) + } + var tp wire.TransportParameters + if err := tp.UnmarshalFromSessionTicket(b); err != nil { + return fmt.Errorf("unmarshaling transport parameters from session ticket failed: %s", err.Error()) + } + t.Parameters = &tp + return nil +} + +const extraPrefix = "quic-go1" + +func addSessionStateExtraPrefix(b []byte) []byte { + return append([]byte(extraPrefix), b...) +} + +func findSessionStateExtraData(extras [][]byte) []byte { + prefix := []byte(extraPrefix) + for _, extra := range extras { + if data, ok := bytes.CutPrefix(extra, prefix); ok { + return data + } + } + return nil +} diff --git a/third_party/quic-go/internal/handshake/session_ticket_test.go b/third_party/quic-go/internal/handshake/session_ticket_test.go new file mode 100644 index 0000000..1585147 --- /dev/null +++ b/third_party/quic-go/internal/handshake/session_ticket_test.go @@ -0,0 +1,46 @@ +package handshake + +import ( + "testing" + + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +func TestMarshalUnmarshalSessionTicket(t *testing.T) { + ticket := &sessionTicket{ + Parameters: &wire.TransportParameters{ + InitialMaxStreamDataBidiLocal: 1, + InitialMaxStreamDataBidiRemote: 2, + ActiveConnectionIDLimit: 10, + MaxDatagramFrameSize: 20, + }, + } + var t2 sessionTicket + require.NoError(t, t2.Unmarshal(ticket.Marshal())) + require.EqualValues(t, 1, t2.Parameters.InitialMaxStreamDataBidiLocal) + require.EqualValues(t, 2, t2.Parameters.InitialMaxStreamDataBidiRemote) + require.EqualValues(t, 10, t2.Parameters.ActiveConnectionIDLimit) + require.EqualValues(t, 20, t2.Parameters.MaxDatagramFrameSize) +} + +func TestUnmarshalRefusesTooShortTicket(t *testing.T) { + err := (&sessionTicket{}).Unmarshal([]byte{}) + require.EqualError(t, err, "failed to read session ticket revision") +} + +func TestUnmarshalRefusesUnknownRevision(t *testing.T) { + b := quicvarint.Append(nil, 1337) + err := (&sessionTicket{}).Unmarshal(b) + require.EqualError(t, err, "unknown session ticket revision: 1337") +} + +func TestUnmarshal0RTTRefusesInvalidTransportParameters(t *testing.T) { + b := quicvarint.Append(nil, sessionTicketRevision) + b = append(b, []byte("foobar")...) + err := (&sessionTicket{}).Unmarshal(b) + require.Error(t, err) + require.Contains(t, err.Error(), "unmarshaling transport parameters from session ticket failed") +} diff --git a/third_party/quic-go/internal/handshake/tls_config_go126.go b/third_party/quic-go/internal/handshake/tls_config_go126.go new file mode 100644 index 0000000..693a08d --- /dev/null +++ b/third_party/quic-go/internal/handshake/tls_config_go126.go @@ -0,0 +1,54 @@ +//go:build !go1.27 + +package handshake + +import ( + "crypto/tls" + "net" +) + +func setupConfigForClient(conf *tls.Config) *tls.Config { + conf = conf.Clone() + conf.MinVersion = tls.VersionTLS13 + return conf +} + +func setupConfigForServer(conf *tls.Config, localAddr, remoteAddr net.Addr) *tls.Config { + // Workaround for https://github.com/golang/go/issues/60506. + // This initializes the session tickets _before_ cloning the config. + _, _ = conf.DecryptTicket(nil, tls.ConnectionState{}) + + conf = conf.Clone() + conf.MinVersion = tls.VersionTLS13 + + // The tls.Config contains two callbacks that pass in a tls.ClientHelloInfo. + // Since crypto/tls doesn't do it, we need to make sure to set the Conn field with a fake net.Conn + // that allows the caller to get the local and the remote address. + if conf.GetConfigForClient != nil { + gcfc := conf.GetConfigForClient + conf.GetConfigForClient = func(info *tls.ClientHelloInfo) (*tls.Config, error) { + info.Conn = &conn{localAddr: localAddr, remoteAddr: remoteAddr} + c, err := gcfc(info) + if c != nil { + // we're returning a tls.Config here, so we need to apply this recursively + c = setupConfigForServer(c, localAddr, remoteAddr) + } + return c, err + } + } + if conf.GetCertificate != nil { + gc := conf.GetCertificate + conf.GetCertificate = func(info *tls.ClientHelloInfo) (*tls.Certificate, error) { + info.Conn = &conn{localAddr: localAddr, remoteAddr: remoteAddr} + return gc(info) + } + } + return conf +} + +func getQUICConfig(tlsConf *tls.Config, _, _ net.Addr) *tls.QUICConfig { + return &tls.QUICConfig{ + TLSConfig: tlsConf, + EnableSessionEvents: true, + } +} diff --git a/third_party/quic-go/internal/handshake/tls_config_go126_test.go b/third_party/quic-go/internal/handshake/tls_config_go126_test.go new file mode 100644 index 0000000..b9303f2 --- /dev/null +++ b/third_party/quic-go/internal/handshake/tls_config_go126_test.go @@ -0,0 +1,99 @@ +//go:build !go1.27 + +package handshake + +import ( + "crypto/tls" + "net" + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestMinimumTLSVersion(t *testing.T) { + local := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42} + remote := &net.UDPAddr{IP: net.IPv4(192, 168, 0, 1), Port: 1337} + + orig := &tls.Config{MinVersion: tls.VersionTLS12} + conf := setupConfigForServer(orig, local, remote) + require.EqualValues(t, tls.VersionTLS13, conf.MinVersion) + // check that the original config wasn't modified + require.EqualValues(t, tls.VersionTLS12, orig.MinVersion) +} + +func TestServerConfigGetCertificate(t *testing.T) { + local := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42} + remote := &net.UDPAddr{IP: net.IPv4(192, 168, 0, 1), Port: 1337} + + var localAddr, remoteAddr net.Addr + tlsConf := &tls.Config{ + GetCertificate: func(info *tls.ClientHelloInfo) (*tls.Certificate, error) { + localAddr = info.Conn.LocalAddr() + remoteAddr = info.Conn.RemoteAddr() + return &tls.Certificate{}, nil + }, + } + conf := setupConfigForServer(tlsConf, local, remote) + _, err := conf.GetCertificate(&tls.ClientHelloInfo{}) + require.NoError(t, err) + require.Equal(t, local, localAddr) + require.Equal(t, remote, remoteAddr) +} + +func TestServerConfigGetConfigForClient(t *testing.T) { + local := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42} + remote := &net.UDPAddr{IP: net.IPv4(192, 168, 0, 1), Port: 1337} + + var localAddr, remoteAddr net.Addr + tlsConf := setupConfigForServer( + &tls.Config{ + GetConfigForClient: func(info *tls.ClientHelloInfo) (*tls.Config, error) { + localAddr = info.Conn.LocalAddr() + remoteAddr = info.Conn.RemoteAddr() + return &tls.Config{}, nil + }, + }, + local, + remote, + ) + conf, err := tlsConf.GetConfigForClient(&tls.ClientHelloInfo{}) + require.NoError(t, err) + require.Equal(t, local, localAddr) + require.Equal(t, remote, remoteAddr) + require.NotNil(t, conf) + require.EqualValues(t, tls.VersionTLS13, conf.MinVersion) +} + +func TestServerConfigGetConfigForClientRecursively(t *testing.T) { + local := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42} + remote := &net.UDPAddr{IP: net.IPv4(192, 168, 0, 1), Port: 1337} + + var localAddr, remoteAddr net.Addr + tlsConf := &tls.Config{} + var innerConf *tls.Config + getCert := func(info *tls.ClientHelloInfo) (*tls.Certificate, error) { + localAddr = info.Conn.LocalAddr() + remoteAddr = info.Conn.RemoteAddr() + return &tls.Certificate{}, nil + } + tlsConf.GetConfigForClient = func(info *tls.ClientHelloInfo) (*tls.Config, error) { + innerConf = tlsConf.Clone() + // set the MaxVersion, so we can check that quic-go doesn't overwrite the user's config + innerConf.MaxVersion = tls.VersionTLS12 + innerConf.GetCertificate = getCert + return innerConf, nil + } + tlsConf = setupConfigForServer(tlsConf, local, remote) + conf, err := tlsConf.GetConfigForClient(&tls.ClientHelloInfo{}) + require.NoError(t, err) + require.NotNil(t, conf) + require.EqualValues(t, tls.VersionTLS13, conf.MinVersion) + _, err = conf.GetCertificate(&tls.ClientHelloInfo{}) + require.NoError(t, err) + require.Equal(t, local, localAddr) + require.Equal(t, remote, remoteAddr) + // make sure that the tls.Config returned by GetConfigForClient isn't modified + require.True(t, reflect.ValueOf(innerConf.GetCertificate).Pointer() == reflect.ValueOf(getCert).Pointer()) + require.EqualValues(t, tls.VersionTLS12, innerConf.MaxVersion) +} diff --git a/third_party/quic-go/internal/handshake/tls_config_go127.go b/third_party/quic-go/internal/handshake/tls_config_go127.go new file mode 100644 index 0000000..7b9d481 --- /dev/null +++ b/third_party/quic-go/internal/handshake/tls_config_go127.go @@ -0,0 +1,24 @@ +//go:build go1.27 + +package handshake + +import ( + "crypto/tls" + "net" +) + +func setupConfigForClient(conf *tls.Config) *tls.Config { + return conf +} + +func setupConfigForServer(conf *tls.Config, _, _ net.Addr) *tls.Config { + return conf +} + +func getQUICConfig(tlsConf *tls.Config, localAddr, remoteAddr net.Addr) *tls.QUICConfig { + return &tls.QUICConfig{ + TLSConfig: tlsConf, + EnableSessionEvents: true, + ClientHelloInfoConn: &conn{localAddr: localAddr, remoteAddr: remoteAddr}, + } +} diff --git a/third_party/quic-go/internal/handshake/tls_conn.go b/third_party/quic-go/internal/handshake/tls_conn.go new file mode 100644 index 0000000..e10fff1 --- /dev/null +++ b/third_party/quic-go/internal/handshake/tls_conn.go @@ -0,0 +1,27 @@ +package handshake + +import ( + "context" + "crypto/tls" +) + +// tlsQUICConn abstracts the QUIC-TLS connection driven by cryptoSetup, so the +// client can be backed either by crypto/tls or, when parroting Chrome, by uTLS. +// +// crypto/tls's *tls.QUICConn satisfies this natively; uTLS is adapted in +// tls_conn_utls.go. Events are normalized to crypto/tls types, which works for +// every field except SessionState, whose internals are unexported and therefore +// not convertible between the two libraries. That is why the uTLS path runs with +// session resumption disabled; see newUTLSQUICClient. +type tlsQUICConn interface { + Start(context.Context) error + NextEvent() tls.QUICEvent + HandleData(tls.QUICEncryptionLevel, []byte) error + SetTransportParameters([]byte) + SendSessionTicket(tls.QUICSessionTicketOptions) error + StoreSession(*tls.SessionState) error + ConnectionState() tls.ConnectionState + Close() error +} + +var _ tlsQUICConn = (*tls.QUICConn)(nil) diff --git a/third_party/quic-go/internal/handshake/tls_conn_utls.go b/third_party/quic-go/internal/handshake/tls_conn_utls.go new file mode 100644 index 0000000..ff8b922 --- /dev/null +++ b/third_party/quic-go/internal/handshake/tls_conn_utls.go @@ -0,0 +1,266 @@ +package handshake + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + + utls "github.com/refraction-networking/utls" + + "github.com/apernet/quic-go/quicvarint" +) + +// utlsQUICConn adapts uTLS's UQUICConn to the tlsQUICConn interface, translating +// its QUIC event and connection-state types back into the crypto/tls +// equivalents. uTLS is a fork of crypto/tls, so the translation is mechanical. +type utlsQUICConn struct { + conn *utls.UQUICConn + // spec is retained because the transport parameters have to be written into + // the ClientHello extension rather than handed to uTLS directly; see + // SetTransportParameters. + spec *utls.ClientHelloSpec +} + +var _ tlsQUICConn = (*utlsQUICConn)(nil) + +// newUTLSQUICClient creates a QUIC-TLS client emitting the parroted ClientHello. +// +// Session resumption and 0-RTT are disabled: crypto/tls and uTLS each have their +// own SessionState with unexported internals, so a uTLS session cannot be +// converted into the *tls.SessionState the resumption path expects. The events +// are turned off at the source rather than half-supported. +func newUTLSQUICClient(tlsConf *tls.Config) (*utlsQUICConn, error) { + uConf, err := utlsConfigFromStd(tlsConf) + if err != nil { + return nil, err + } + spec := chromeQUICClientHelloSpec(tlsConf.NextProtos) + + conn := utls.UQUICClient(&utls.QUICConfig{TLSConfig: uConf}, utls.HelloCustom) + if err := conn.ApplyPreset(spec); err != nil { + return nil, fmt.Errorf("applying Chrome ClientHello spec: %w", err) + } + return &utlsQUICConn{conn: conn, spec: spec}, nil +} + +// utlsConfigFromStd converts a crypto/tls client config into the uTLS +// equivalent. +// +// uTLS ships no converter, so this copies field by field. Any field that cannot +// be carried across is a hard error rather than a silent omission: dropping +// something like VerifyConnection would quietly weaken certificate validation. +func utlsConfigFromStd(c *tls.Config) (*utls.Config, error) { + if c == nil { + return &utls.Config{MinVersion: utls.VersionTLS13}, nil + } + if c.VerifyConnection != nil { + // Takes a crypto/tls ConnectionState, which uTLS will never produce. + return nil, errors.New("quic: tls.Config.VerifyConnection is not supported with ChromeParrot") + } + if c.GetConfigForClient != nil || len(c.Certificates) > 0 || c.GetCertificate != nil { + return nil, errors.New("quic: server-side tls.Config fields are not supported with ChromeParrot") + } + + uc := &utls.Config{ + Rand: c.Rand, + Time: c.Time, + RootCAs: c.RootCAs, + NextProtos: c.NextProtos, + ServerName: c.ServerName, + InsecureSkipVerify: c.InsecureSkipVerify, + VerifyPeerCertificate: c.VerifyPeerCertificate, + KeyLogWriter: c.KeyLogWriter, + // TLS 1.3 only, which QUIC requires regardless. + MinVersion: utls.VersionTLS13, + MaxVersion: utls.VersionTLS13, + // Resumption is off; see newUTLSQUICClient. + SessionTicketsDisabled: true, + } + + if c.GetClientCertificate != nil { + get := c.GetClientCertificate + uc.GetClientCertificate = func(cri *utls.CertificateRequestInfo) (*utls.Certificate, error) { + cert, err := get(&tls.CertificateRequestInfo{ + AcceptableCAs: cri.AcceptableCAs, + SignatureSchemes: signatureSchemesToStd(cri.SignatureSchemes), + Version: cri.Version, + }) + if err != nil { + return nil, err + } + if cert == nil { + return &utls.Certificate{}, nil + } + return &utls.Certificate{ + Certificate: cert.Certificate, + PrivateKey: cert.PrivateKey, + OCSPStaple: cert.OCSPStaple, + SignedCertificateTimestamps: cert.SignedCertificateTimestamps, + Leaf: cert.Leaf, + }, nil + } + } + + // GREASE ECH is covered by the ClientHello spec; a caller-supplied config list + // takes precedence. + if len(c.EncryptedClientHelloConfigList) > 0 { + uc.EncryptedClientHelloConfigList = c.EncryptedClientHelloConfigList + } + return uc, nil +} + +func signatureSchemesToStd(in []utls.SignatureScheme) []tls.SignatureScheme { + out := make([]tls.SignatureScheme, len(in)) + for i, s := range in { + out[i] = tls.SignatureScheme(s) + } + return out +} + +func (c *utlsQUICConn) Start(ctx context.Context) error { return c.conn.Start(ctx) } +func (c *utlsQUICConn) Close() error { return c.conn.Close() } + +func (c *utlsQUICConn) HandleData(level tls.QUICEncryptionLevel, data []byte) error { + return c.conn.HandleData(utlsEncryptionLevel(level), data) +} + +// SetTransportParameters installs quic-go's marshalled transport parameters into +// the ClientHello. +// +// uTLS's own SetTransportParameters does not reach the ClientHello when a preset +// is in use, so the bytes must be written into the spec's +// quic_transport_parameters extension. uTLS models that extension as (id, value) +// pairs and marshals it itself, so the blob is split back into pairs here. +// Splitting rather than re-deriving preserves our per-connection ordering. +func (c *utlsQUICConn) SetTransportParameters(params []byte) { + // Still call through so uTLS's internal copy stays consistent. + c.conn.SetTransportParameters(params) + + tps, err := splitTransportParameters(params) + if err != nil { + // Our own marshaller produced these, so this is unreachable short of a bug. + panic(fmt.Sprintf("handshake BUG: cannot split marshalled transport parameters: %s", err)) + } + for _, ext := range c.spec.Extensions { + if qtp, ok := ext.(*utls.QUICTransportParametersExtension); ok { + qtp.TransportParameters = tps + return + } + } + panic("handshake BUG: Chrome ClientHello spec has no quic_transport_parameters extension") +} + +// splitTransportParameters parses a marshalled transport parameter blob back into +// individual (id, value) pairs, preserving order. +func splitTransportParameters(b []byte) (utls.TransportParameters, error) { + var tps utls.TransportParameters + for len(b) > 0 { + id, n, err := quicvarint.Parse(b) + if err != nil { + return nil, err + } + b = b[n:] + l, n, err := quicvarint.Parse(b) + if err != nil { + return nil, err + } + b = b[n:] + if uint64(len(b)) < l { + return nil, fmt.Errorf("transport parameter 0x%x truncated: want %d bytes, have %d", id, l, len(b)) + } + val := make([]byte, l) + copy(val, b[:l]) + b = b[l:] + tps = append(tps, &utls.FakeQUICTransportParameter{Id: id, Val: val}) + } + return tps, nil +} + +func (c *utlsQUICConn) NextEvent() tls.QUICEvent { + ev := c.conn.NextEvent() + out := tls.QUICEvent{ + Level: stdEncryptionLevel(ev.Level), + Data: ev.Data, + Suite: ev.Suite, + } + switch ev.Kind { + case utls.QUICNoEvent: + out.Kind = tls.QUICNoEvent + case utls.QUICSetReadSecret: + out.Kind = tls.QUICSetReadSecret + case utls.QUICSetWriteSecret: + out.Kind = tls.QUICSetWriteSecret + case utls.QUICWriteData: + out.Kind = tls.QUICWriteData + case utls.QUICTransportParameters: + out.Kind = tls.QUICTransportParameters + case utls.QUICTransportParametersRequired: + out.Kind = tls.QUICTransportParametersRequired + case utls.QUICRejectedEarlyData: + out.Kind = tls.QUICRejectedEarlyData + case utls.QUICHandshakeDone: + out.Kind = tls.QUICHandshakeDone + default: + // QUICStoreSession and QUICResumeSession carry a *utls.SessionState that + // cannot be converted. They only fire when EnableSessionEvents is set, + // which newUTLSQUICClient never does, so reaching this is a bug. + panic(fmt.Sprintf("handshake BUG: unexpected uTLS QUIC event kind %d", ev.Kind)) + } + return out +} + +func (c *utlsQUICConn) SendSessionTicket(tls.QUICSessionTicketOptions) error { + return errors.New("quic: SendSessionTicket is server-only and unsupported with ChromeParrot") +} + +func (c *utlsQUICConn) StoreSession(*tls.SessionState) error { + return errors.New("quic: session resumption is unsupported with ChromeParrot") +} + +func (c *utlsQUICConn) ConnectionState() tls.ConnectionState { + s := c.conn.ConnectionState() + return tls.ConnectionState{ + Version: s.Version, + HandshakeComplete: s.HandshakeComplete, + DidResume: s.DidResume, + CipherSuite: s.CipherSuite, + NegotiatedProtocol: s.NegotiatedProtocol, + NegotiatedProtocolIsMutual: true, + ServerName: s.ServerName, + PeerCertificates: s.PeerCertificates, + VerifiedChains: s.VerifiedChains, + SignedCertificateTimestamps: s.SignedCertificateTimestamps, + OCSPResponse: s.OCSPResponse, + } +} + +func utlsEncryptionLevel(l tls.QUICEncryptionLevel) utls.QUICEncryptionLevel { + switch l { + case tls.QUICEncryptionLevelInitial: + return utls.QUICEncryptionLevelInitial + case tls.QUICEncryptionLevelEarly: + return utls.QUICEncryptionLevelEarly + case tls.QUICEncryptionLevelHandshake: + return utls.QUICEncryptionLevelHandshake + case tls.QUICEncryptionLevelApplication: + return utls.QUICEncryptionLevelApplication + default: + panic(fmt.Sprintf("handshake BUG: unknown encryption level %d", l)) + } +} + +func stdEncryptionLevel(l utls.QUICEncryptionLevel) tls.QUICEncryptionLevel { + switch l { + case utls.QUICEncryptionLevelInitial: + return tls.QUICEncryptionLevelInitial + case utls.QUICEncryptionLevelEarly: + return tls.QUICEncryptionLevelEarly + case utls.QUICEncryptionLevelHandshake: + return tls.QUICEncryptionLevelHandshake + case utls.QUICEncryptionLevelApplication: + return tls.QUICEncryptionLevelApplication + default: + panic(fmt.Sprintf("handshake BUG: unknown uTLS encryption level %d", l)) + } +} diff --git a/third_party/quic-go/internal/handshake/token_generator.go b/third_party/quic-go/internal/handshake/token_generator.go new file mode 100644 index 0000000..092056f --- /dev/null +++ b/third_party/quic-go/internal/handshake/token_generator.go @@ -0,0 +1,126 @@ +package handshake + +import ( + "bytes" + "encoding/asn1" + "fmt" + "net" + "time" + + "github.com/apernet/quic-go/internal/protocol" +) + +const ( + tokenPrefixIP byte = iota + tokenPrefixString +) + +// A Token is derived from the client address and can be used to verify the ownership of this address. +type Token struct { + IsRetryToken bool + SentTime time.Time + encodedRemoteAddr []byte + // only set for tokens sent in NEW_TOKEN frames + RTT time.Duration + // only set for retry tokens + OriginalDestConnectionID protocol.ConnectionID + RetrySrcConnectionID protocol.ConnectionID +} + +// ValidateRemoteAddr validates the address, but does not check expiration +func (t *Token) ValidateRemoteAddr(addr net.Addr) bool { + return bytes.Equal(encodeRemoteAddr(addr), t.encodedRemoteAddr) +} + +// token is the struct that is used for ASN1 serialization and deserialization +type token struct { + IsRetryToken bool + RemoteAddr []byte + Timestamp int64 + RTT int64 // in mus + OriginalDestConnectionID []byte + RetrySrcConnectionID []byte +} + +// A TokenGenerator generates tokens +type TokenGenerator struct { + tokenProtector tokenProtector +} + +// NewTokenGenerator initializes a new TokenGenerator +func NewTokenGenerator(key TokenProtectorKey) *TokenGenerator { + return &TokenGenerator{tokenProtector: *newTokenProtector(key)} +} + +// NewRetryToken generates a new token for a Retry for a given source address +func (g *TokenGenerator) NewRetryToken( + raddr net.Addr, + origDestConnID protocol.ConnectionID, + retrySrcConnID protocol.ConnectionID, +) ([]byte, error) { + data, err := asn1.Marshal(token{ + IsRetryToken: true, + RemoteAddr: encodeRemoteAddr(raddr), + OriginalDestConnectionID: origDestConnID.Bytes(), + RetrySrcConnectionID: retrySrcConnID.Bytes(), + Timestamp: time.Now().UnixNano(), + }) + if err != nil { + return nil, err + } + return g.tokenProtector.NewToken(data) +} + +// NewToken generates a new token to be sent in a NEW_TOKEN frame +func (g *TokenGenerator) NewToken(raddr net.Addr, rtt time.Duration) ([]byte, error) { + data, err := asn1.Marshal(token{ + RemoteAddr: encodeRemoteAddr(raddr), + Timestamp: time.Now().UnixNano(), + RTT: rtt.Microseconds(), + }) + if err != nil { + return nil, err + } + return g.tokenProtector.NewToken(data) +} + +// DecodeToken decodes a token +func (g *TokenGenerator) DecodeToken(encrypted []byte) (*Token, error) { + // if the client didn't send any token, DecodeToken will be called with a nil-slice + if len(encrypted) == 0 { + return nil, nil + } + + data, err := g.tokenProtector.DecodeToken(encrypted) + if err != nil { + return nil, err + } + t := &token{} + rest, err := asn1.Unmarshal(data, t) + if err != nil { + return nil, err + } + if len(rest) != 0 { + return nil, fmt.Errorf("rest when unpacking token: %d", len(rest)) + } + token := &Token{ + IsRetryToken: t.IsRetryToken, + SentTime: time.Unix(0, t.Timestamp), + encodedRemoteAddr: t.RemoteAddr, + } + if t.IsRetryToken { + token.OriginalDestConnectionID = protocol.ParseConnectionID(t.OriginalDestConnectionID) + token.RetrySrcConnectionID = protocol.ParseConnectionID(t.RetrySrcConnectionID) + } else { + token.RTT = time.Duration(t.RTT) * time.Microsecond + } + return token, nil +} + +// encodeRemoteAddr encodes a remote address such that it can be saved in the token +func encodeRemoteAddr(remoteAddr net.Addr) []byte { + if udpAddr, ok := remoteAddr.(*net.UDPAddr); ok { + return append([]byte{tokenPrefixIP}, udpAddr.IP...) + } + return append([]byte{tokenPrefixString}, []byte(remoteAddr.String())...) +} diff --git a/third_party/quic-go/internal/handshake/token_generator_test.go b/third_party/quic-go/internal/handshake/token_generator_test.go new file mode 100644 index 0000000..421bfe4 --- /dev/null +++ b/third_party/quic-go/internal/handshake/token_generator_test.go @@ -0,0 +1,140 @@ +package handshake + +import ( + "crypto/rand" + "encoding/asn1" + "net" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func newTokenGenerator(t *testing.T) *TokenGenerator { + var key TokenProtectorKey + _, err := rand.Read(key[:]) + require.NoError(t, err) + return NewTokenGenerator(key) +} + +func TestTokenGeneratorNilTokens(t *testing.T) { + tokenGen := newTokenGenerator(t) + nilToken, err := tokenGen.DecodeToken(nil) + require.NoError(t, err) + require.Nil(t, nilToken) +} + +func TestTokenGeneratorValidToken(t *testing.T) { + tokenGen := newTokenGenerator(t) + + addr := &net.UDPAddr{IP: net.IPv4(192, 168, 0, 1), Port: 1337} + connID1 := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}) + connID2 := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}) + tokenEnc, err := tokenGen.NewRetryToken(addr, connID1, connID2) + require.NoError(t, err) + decodedToken, err := tokenGen.DecodeToken(tokenEnc) + require.NoError(t, err) + require.True(t, decodedToken.ValidateRemoteAddr(addr)) + require.False(t, decodedToken.ValidateRemoteAddr(&net.UDPAddr{IP: net.IPv4(192, 168, 0, 2), Port: 1337})) + require.WithinDuration(t, time.Now(), decodedToken.SentTime, 100*time.Millisecond) + require.Equal(t, connID1, decodedToken.OriginalDestConnectionID) + require.Equal(t, connID2, decodedToken.RetrySrcConnectionID) +} + +func TestTokenGeneratorRejectsInvalidTokens(t *testing.T) { + tokenGen := newTokenGenerator(t) + + _, err := tokenGen.DecodeToken([]byte("invalid token")) + require.Error(t, err) + require.Contains(t, err.Error(), "too short") +} + +func TestTokenGeneratorDecodingFailed(t *testing.T) { + tokenGen := newTokenGenerator(t) + + invalidToken, err := tokenGen.tokenProtector.NewToken([]byte("foobar")) + require.NoError(t, err) + _, err = tokenGen.DecodeToken(invalidToken) + require.Error(t, err) + require.Contains(t, err.Error(), "asn1") +} + +func TestTokenGeneratorAdditionalPayload(t *testing.T) { + tokenGen := newTokenGenerator(t) + + tok, err := asn1.Marshal(token{RemoteAddr: []byte("foobar")}) + require.NoError(t, err) + tok = append(tok, []byte("rest")...) + enc, err := tokenGen.tokenProtector.NewToken(tok) + require.NoError(t, err) + _, err = tokenGen.DecodeToken(enc) + require.EqualError(t, err, "rest when unpacking token: 4") +} + +func TestTokenGeneratorEmptyTokens(t *testing.T) { + tokenGen := newTokenGenerator(t) + + emptyTok, err := asn1.Marshal(token{RemoteAddr: []byte("")}) + require.NoError(t, err) + emptyEnc, err := tokenGen.tokenProtector.NewToken(emptyTok) + require.NoError(t, err) + _, err = tokenGen.DecodeToken(emptyEnc) + require.NoError(t, err) +} + +func TestTokenGeneratorIPv6(t *testing.T) { + tokenGen := newTokenGenerator(t) + + addresses := []string{ + "2001:db8::68", + "2001:0000:4136:e378:8000:63bf:3fff:fdd2", + "2001::1", + "ff01:0:0:0:0:0:0:2", + } + for _, addr := range addresses { + ip := net.ParseIP(addr) + require.NotNil(t, ip) + raddr := &net.UDPAddr{IP: ip, Port: 1337} + tokenEnc, err := tokenGen.NewRetryToken(raddr, protocol.ConnectionID{}, protocol.ConnectionID{}) + require.NoError(t, err) + token, err := tokenGen.DecodeToken(tokenEnc) + require.NoError(t, err) + require.True(t, token.ValidateRemoteAddr(raddr)) + require.WithinDuration(t, time.Now(), token.SentTime, 100*time.Millisecond) + } +} + +func TestTokenGeneratorNonUDPAddr(t *testing.T) { + tokenGen := newTokenGenerator(t) + + raddr := &net.TCPAddr{IP: net.IPv4(192, 168, 13, 37), Port: 1337} + tokenEnc, err := tokenGen.NewRetryToken(raddr, protocol.ConnectionID{}, protocol.ConnectionID{}) + require.NoError(t, err) + token, err := tokenGen.DecodeToken(tokenEnc) + require.NoError(t, err) + require.True(t, token.ValidateRemoteAddr(raddr)) + require.False(t, token.ValidateRemoteAddr(&net.TCPAddr{IP: net.IPv4(192, 168, 13, 37), Port: 1338})) + require.WithinDuration(t, time.Now(), token.SentTime, 100*time.Millisecond) +} + +func BenchmarkTokenGeneratorDecodeToken(b *testing.B) { + b.ReportAllocs() + + var key TokenProtectorKey + _, err := rand.Read(key[:]) + require.NoError(b, err) + tokenGen := NewTokenGenerator(key) + addr := &net.UDPAddr{IP: net.IPv4(192, 168, 0, 1), Port: 1337} + connID1 := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}) + connID2 := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}) + tokenEnc, err := tokenGen.NewRetryToken(addr, connID1, connID2) + require.NoError(b, err) + + for b.Loop() { + if _, err := tokenGen.DecodeToken(tokenEnc); err != nil { + b.Fatal(err) + } + } +} diff --git a/third_party/quic-go/internal/handshake/token_protector.go b/third_party/quic-go/internal/handshake/token_protector.go new file mode 100644 index 0000000..74cd7d7 --- /dev/null +++ b/third_party/quic-go/internal/handshake/token_protector.go @@ -0,0 +1,78 @@ +package handshake + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/hkdf" + "crypto/rand" + "crypto/sha256" + "fmt" +) + +// TokenProtectorKey is the key used to encrypt both Retry and session resumption tokens. +type TokenProtectorKey [32]byte + +const tokenSaltSize = 32 + +// tokenProtector is used to create and verify a token +type tokenProtector struct { + key TokenProtectorKey +} + +// newTokenProtector creates a source for source address tokens +func newTokenProtector(key TokenProtectorKey) *tokenProtector { + return &tokenProtector{key: key} +} + +// NewToken encodes data into a new token. +func (s *tokenProtector) NewToken(data []byte) ([]byte, error) { + var salt [tokenSaltSize]byte + if _, err := rand.Read(salt[:]); err != nil { + return nil, err + } + aead, err := s.createAEAD(salt[:]) + if err != nil { + return nil, err + } + return append(salt[:], aead.Seal(nil, nil, data, nil)...), nil +} + +// DecodeToken decodes a token. +func (s *tokenProtector) DecodeToken(p []byte) ([]byte, error) { + if len(p) < tokenSaltSize { + return nil, fmt.Errorf("token too short: %d", len(p)) + } + salt := p[:tokenSaltSize] + aead, err := s.createAEAD(salt) + if err != nil { + return nil, err + } + if len(p[tokenSaltSize:]) < aead.Overhead() { + return nil, fmt.Errorf("token too short: %d", len(p)) + } + return aead.Open(nil, nil, p[tokenSaltSize:], nil) +} + +const tokenProtectorHKDFInfo = "quic-go token source" + +func (s *tokenProtector) createAEAD(salt []byte) (cipher.AEAD, error) { + prk, err := hkdf.Extract(sha256.New, s.key[:], salt) + if err != nil { + return nil, err + } + + key, err := hkdf.Expand(sha256.New, prk, tokenProtectorHKDFInfo, 32) + if err != nil { + return nil, err + } + + c, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + aead, err := cipher.NewGCMWithRandomNonce(c) + if err != nil { + return nil, err + } + return aead, nil +} diff --git a/third_party/quic-go/internal/handshake/token_protector_test.go b/third_party/quic-go/internal/handshake/token_protector_test.go new file mode 100644 index 0000000..f2a290a --- /dev/null +++ b/third_party/quic-go/internal/handshake/token_protector_test.go @@ -0,0 +1,70 @@ +package handshake + +import ( + "crypto/rand" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestTokenProtectorEncodeAndDecode(t *testing.T) { + var key TokenProtectorKey + rand.Read(key[:]) + tp := newTokenProtector(key) + + token, err := tp.NewToken([]byte("foobar")) + require.NoError(t, err) + require.NotContains(t, string(token), "foobar") + + decoded, err := tp.DecodeToken(token) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), decoded) +} + +func TestTokenProtectorDifferentKeys(t *testing.T) { + var key1, key2 TokenProtectorKey + rand.Read(key1[:]) + rand.Read(key2[:]) + tp1 := newTokenProtector(key1) + tp2 := newTokenProtector(key2) + + t1, err := tp1.NewToken([]byte("foo")) + require.NoError(t, err) + t2, err := tp2.NewToken([]byte("foo")) + require.NoError(t, err) + + _, err = tp1.DecodeToken(t1) + require.NoError(t, err) + _, err = tp1.DecodeToken(t2) + require.Error(t, err) + + tp3 := newTokenProtector(key1) + _, err = tp3.DecodeToken(t1) + require.NoError(t, err) + _, err = tp3.DecodeToken(t2) + require.Error(t, err) +} + +func TestTokenProtectorInvalidTokens(t *testing.T) { + var key TokenProtectorKey + rand.Read(key[:]) + tp := newTokenProtector(key) + + token, err := tp.NewToken([]byte("foobar")) + require.NoError(t, err) + _, err = tp.DecodeToken(token[1:]) + require.Error(t, err) + require.Contains(t, err.Error(), "message authentication failed") +} + +func TestTokenProtectorTooShortTokens(t *testing.T) { + var key TokenProtectorKey + rand.Read(key[:]) + tp := newTokenProtector(key) + + _, err := tp.DecodeToken([]byte("foobar")) + require.EqualError(t, err, "token too short: 6") + + _, err = tp.DecodeToken(make([]byte, tokenSaltSize)) + require.EqualError(t, err, "token too short: 32") +} diff --git a/third_party/quic-go/internal/handshake/updatable_aead.go b/third_party/quic-go/internal/handshake/updatable_aead.go new file mode 100644 index 0000000..38dcefd --- /dev/null +++ b/third_party/quic-go/internal/handshake/updatable_aead.go @@ -0,0 +1,372 @@ +package handshake + +import ( + "crypto" + "crypto/cipher" + "crypto/tls" + "encoding/binary" + "fmt" + "sync/atomic" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +var keyUpdateInterval atomic.Uint64 + +func init() { + keyUpdateInterval.Store(protocol.KeyUpdateInterval) +} + +func SetKeyUpdateInterval(v uint64) (reset func()) { + old := keyUpdateInterval.Swap(v) + return func() { keyUpdateInterval.Store(old) } +} + +// FirstKeyUpdateInterval is the maximum number of packets we send or receive before initiating the first key update. +// It's a package-level variable to allow modifying it for testing purposes. +var FirstKeyUpdateInterval uint64 = 100 + +type updatableAEAD struct { + suite cipherSuite + + keyPhase protocol.KeyPhase + largestAcked protocol.PacketNumber + firstPacketNumber protocol.PacketNumber + handshakeConfirmed bool + + invalidPacketLimit uint64 + invalidPacketCount uint64 + + // Time when the keys should be dropped. Keys are dropped on the next call to Open(). + prevRcvAEADExpiry monotime.Time + prevRcvAEAD cipher.AEAD + + firstRcvdWithCurrentKey protocol.PacketNumber + firstSentWithCurrentKey protocol.PacketNumber + highestRcvdPN protocol.PacketNumber // highest packet number received (which could be successfully unprotected) + numRcvdWithCurrentKey uint64 + numSentWithCurrentKey uint64 + rcvAEAD cipher.AEAD + sendAEAD cipher.AEAD + // caches cipher.AEAD.Overhead(). This speeds up calls to Overhead(). + aeadOverhead int + + nextRcvAEAD cipher.AEAD + nextSendAEAD cipher.AEAD + nextRcvTrafficSecret []byte + nextSendTrafficSecret []byte + + headerDecrypter headerProtector + headerEncrypter headerProtector + + rttStats *utils.RTTStats + + qlogger qlogwriter.Recorder + logger utils.Logger + version protocol.Version + + // use a single slice to avoid allocations + nonceBuf []byte +} + +var ( + _ ShortHeaderOpener = &updatableAEAD{} + _ ShortHeaderSealer = &updatableAEAD{} +) + +func newUpdatableAEAD(rttStats *utils.RTTStats, qlogger qlogwriter.Recorder, logger utils.Logger, version protocol.Version) *updatableAEAD { + return &updatableAEAD{ + firstPacketNumber: protocol.InvalidPacketNumber, + largestAcked: protocol.InvalidPacketNumber, + firstRcvdWithCurrentKey: protocol.InvalidPacketNumber, + firstSentWithCurrentKey: protocol.InvalidPacketNumber, + rttStats: rttStats, + qlogger: qlogger, + logger: logger, + version: version, + } +} + +func (a *updatableAEAD) rollKeys() { + if a.prevRcvAEAD != nil { + a.logger.Debugf("Dropping key phase %d ahead of scheduled time. Drop time was: %s", a.keyPhase-1, a.prevRcvAEADExpiry) + if a.qlogger != nil { + a.qlogger.RecordEvent(qlog.KeyDiscarded{ + KeyType: qlog.KeyTypeClient1RTT, + KeyPhase: a.keyPhase - 1, + }) + a.qlogger.RecordEvent(qlog.KeyDiscarded{ + KeyType: qlog.KeyTypeServer1RTT, + KeyPhase: a.keyPhase - 1, + }) + } + a.prevRcvAEADExpiry = 0 + } + + a.keyPhase++ + a.firstRcvdWithCurrentKey = protocol.InvalidPacketNumber + a.firstSentWithCurrentKey = protocol.InvalidPacketNumber + a.numRcvdWithCurrentKey = 0 + a.numSentWithCurrentKey = 0 + a.prevRcvAEAD = a.rcvAEAD + a.rcvAEAD = a.nextRcvAEAD + a.sendAEAD = a.nextSendAEAD + + a.nextRcvTrafficSecret = a.getNextTrafficSecret(a.suite.Hash, a.nextRcvTrafficSecret) + a.nextSendTrafficSecret = a.getNextTrafficSecret(a.suite.Hash, a.nextSendTrafficSecret) + a.nextRcvAEAD = createAEAD(a.suite, a.nextRcvTrafficSecret, a.version) + a.nextSendAEAD = createAEAD(a.suite, a.nextSendTrafficSecret, a.version) +} + +func (a *updatableAEAD) startKeyDropTimer(now monotime.Time) { + d := 3 * a.rttStats.PTO(true) + a.logger.Debugf("Starting key drop timer to drop key phase %d (in %s)", a.keyPhase-1, d) + a.prevRcvAEADExpiry = now.Add(d) +} + +func (a *updatableAEAD) getNextTrafficSecret(hash crypto.Hash, ts []byte) []byte { + return hkdfExpandLabel(hash, ts, []byte{}, "quic ku", hash.Size()) +} + +// SetReadKey sets the read key. +// For the client, this function is called before SetWriteKey. +// For the server, this function is called after SetWriteKey. +func (a *updatableAEAD) SetReadKey(suite cipherSuite, trafficSecret []byte) { + a.rcvAEAD = createAEAD(suite, trafficSecret, a.version) + a.headerDecrypter = newHeaderProtector(suite, trafficSecret, false, a.version) + if a.suite.ID == 0 { // suite is not set yet + a.setAEADParameters(a.rcvAEAD, suite) + } + + a.nextRcvTrafficSecret = a.getNextTrafficSecret(suite.Hash, trafficSecret) + a.nextRcvAEAD = createAEAD(suite, a.nextRcvTrafficSecret, a.version) +} + +// SetWriteKey sets the write key. +// For the client, this function is called after SetReadKey. +// For the server, this function is called before SetReadKey. +func (a *updatableAEAD) SetWriteKey(suite cipherSuite, trafficSecret []byte) { + a.sendAEAD = createAEAD(suite, trafficSecret, a.version) + a.headerEncrypter = newHeaderProtector(suite, trafficSecret, false, a.version) + if a.suite.ID == 0 { // suite is not set yet + a.setAEADParameters(a.sendAEAD, suite) + } + + a.nextSendTrafficSecret = a.getNextTrafficSecret(suite.Hash, trafficSecret) + a.nextSendAEAD = createAEAD(suite, a.nextSendTrafficSecret, a.version) +} + +func (a *updatableAEAD) setAEADParameters(aead cipher.AEAD, suite cipherSuite) { + a.nonceBuf = make([]byte, aead.NonceSize()) + a.aeadOverhead = aead.Overhead() + a.suite = suite + switch suite.ID { + case tls.TLS_AES_128_GCM_SHA256, tls.TLS_AES_256_GCM_SHA384: + a.invalidPacketLimit = protocol.InvalidPacketLimitAES + case tls.TLS_CHACHA20_POLY1305_SHA256: + a.invalidPacketLimit = protocol.InvalidPacketLimitChaCha + default: + panic(fmt.Sprintf("unknown cipher suite %d", suite.ID)) + } +} + +func (a *updatableAEAD) DecodePacketNumber(wirePN protocol.PacketNumber, wirePNLen protocol.PacketNumberLen) protocol.PacketNumber { + return protocol.DecodePacketNumber(wirePNLen, a.highestRcvdPN, wirePN) +} + +func (a *updatableAEAD) Open(dst, src []byte, rcvTime monotime.Time, pn protocol.PacketNumber, kp protocol.KeyPhaseBit, ad []byte) ([]byte, error) { + dec, err := a.open(dst, src, rcvTime, pn, kp, ad) + if err == ErrDecryptionFailed { + a.invalidPacketCount++ + if a.invalidPacketCount >= a.invalidPacketLimit { + return nil, &qerr.TransportError{ErrorCode: qerr.AEADLimitReached} + } + } + if err == nil { + a.highestRcvdPN = max(a.highestRcvdPN, pn) + } + return dec, err +} + +func (a *updatableAEAD) open(dst, src []byte, rcvTime monotime.Time, pn protocol.PacketNumber, kp protocol.KeyPhaseBit, ad []byte) ([]byte, error) { + if a.prevRcvAEAD != nil && !a.prevRcvAEADExpiry.IsZero() && rcvTime.After(a.prevRcvAEADExpiry) { + a.prevRcvAEAD = nil + a.logger.Debugf("Dropping key phase %d", a.keyPhase-1) + a.prevRcvAEADExpiry = 0 + if a.qlogger != nil { + a.qlogger.RecordEvent(qlog.KeyDiscarded{ + KeyType: qlog.KeyTypeClient1RTT, + KeyPhase: a.keyPhase - 1, + }) + a.qlogger.RecordEvent(qlog.KeyDiscarded{ + KeyType: qlog.KeyTypeServer1RTT, + KeyPhase: a.keyPhase - 1, + }) + } + } + binary.BigEndian.PutUint64(a.nonceBuf[len(a.nonceBuf)-8:], uint64(pn)) + if kp != a.keyPhase.Bit() { + if a.keyPhase > 0 && a.firstRcvdWithCurrentKey == protocol.InvalidPacketNumber || pn < a.firstRcvdWithCurrentKey { + if a.prevRcvAEAD == nil { + return nil, ErrKeysDropped + } + // we updated the key, but the peer hasn't updated yet + dec, err := a.prevRcvAEAD.Open(dst, a.nonceBuf, src, ad) + if err != nil { + err = ErrDecryptionFailed + } + return dec, err + } + // try opening the packet with the next key phase + dec, err := a.nextRcvAEAD.Open(dst, a.nonceBuf, src, ad) + if err != nil { + return nil, ErrDecryptionFailed + } + // Opening succeeded. Check if the peer was allowed to update. + if a.keyPhase > 0 && a.firstSentWithCurrentKey == protocol.InvalidPacketNumber { + return nil, &qerr.TransportError{ + ErrorCode: qerr.KeyUpdateError, + ErrorMessage: "keys updated too quickly", + } + } + a.rollKeys() + a.logger.Debugf("Peer updated keys to %d", a.keyPhase) + // The peer initiated this key update. It's safe to drop the keys for the previous generation now. + // Start a timer to drop the previous key generation. + a.startKeyDropTimer(rcvTime) + if a.qlogger != nil { + a.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateRemote, + KeyType: qlog.KeyTypeClient1RTT, + KeyPhase: a.keyPhase, + }) + a.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateRemote, + KeyType: qlog.KeyTypeServer1RTT, + KeyPhase: a.keyPhase, + }) + } + a.firstRcvdWithCurrentKey = pn + return dec, err + } + // The AEAD we're using here will be the qtls.aeadAESGCM13. + // It uses the nonce provided here and XOR it with the IV. + dec, err := a.rcvAEAD.Open(dst, a.nonceBuf, src, ad) + if err != nil { + return dec, ErrDecryptionFailed + } + a.numRcvdWithCurrentKey++ + if a.firstRcvdWithCurrentKey == protocol.InvalidPacketNumber { + // We initiated the key updated, and now we received the first packet protected with the new key phase. + // Therefore, we are certain that the peer rolled its keys as well. Start a timer to drop the old keys. + if a.keyPhase > 0 { + a.logger.Debugf("Peer confirmed key update to phase %d", a.keyPhase) + a.startKeyDropTimer(rcvTime) + } + a.firstRcvdWithCurrentKey = pn + } + return dec, err +} + +func (a *updatableAEAD) Seal(dst, src []byte, pn protocol.PacketNumber, ad []byte) []byte { + if a.firstSentWithCurrentKey == protocol.InvalidPacketNumber { + a.firstSentWithCurrentKey = pn + } + if a.firstPacketNumber == protocol.InvalidPacketNumber { + a.firstPacketNumber = pn + } + a.numSentWithCurrentKey++ + binary.BigEndian.PutUint64(a.nonceBuf[len(a.nonceBuf)-8:], uint64(pn)) + // The AEAD we're using here will be the qtls.aeadAESGCM13. + // It uses the nonce provided here and XOR it with the IV. + return a.sendAEAD.Seal(dst, a.nonceBuf, src, ad) +} + +func (a *updatableAEAD) SetLargestAcked(pn protocol.PacketNumber) error { + if a.firstSentWithCurrentKey != protocol.InvalidPacketNumber && + pn >= a.firstSentWithCurrentKey && a.numRcvdWithCurrentKey == 0 { + return &qerr.TransportError{ + ErrorCode: qerr.KeyUpdateError, + ErrorMessage: fmt.Sprintf("received ACK for key phase %d, but peer didn't update keys", a.keyPhase), + } + } + a.largestAcked = pn + return nil +} + +func (a *updatableAEAD) SetHandshakeConfirmed() { + a.handshakeConfirmed = true +} + +func (a *updatableAEAD) updateAllowed() bool { + if !a.handshakeConfirmed { + return false + } + // the first key update is allowed as soon as the handshake is confirmed + return a.keyPhase == 0 || + // subsequent key updates as soon as a packet sent with that key phase has been acknowledged + (a.firstSentWithCurrentKey != protocol.InvalidPacketNumber && + a.largestAcked != protocol.InvalidPacketNumber && + a.largestAcked >= a.firstSentWithCurrentKey) +} + +func (a *updatableAEAD) shouldInitiateKeyUpdate() bool { + if !a.updateAllowed() { + return false + } + // Initiate the first key update shortly after the handshake, in order to exercise the key update mechanism. + if a.keyPhase == 0 { + if a.numRcvdWithCurrentKey >= FirstKeyUpdateInterval || a.numSentWithCurrentKey >= FirstKeyUpdateInterval { + return true + } + } + if a.numRcvdWithCurrentKey >= keyUpdateInterval.Load() { + a.logger.Debugf("Received %d packets with current key phase. Initiating key update to the next key phase: %d", a.numRcvdWithCurrentKey, a.keyPhase+1) + return true + } + if a.numSentWithCurrentKey >= keyUpdateInterval.Load() { + a.logger.Debugf("Sent %d packets with current key phase. Initiating key update to the next key phase: %d", a.numSentWithCurrentKey, a.keyPhase+1) + return true + } + return false +} + +func (a *updatableAEAD) KeyPhase() protocol.KeyPhaseBit { + if a.shouldInitiateKeyUpdate() { + a.rollKeys() + if a.qlogger != nil { + a.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateLocal, + KeyType: qlog.KeyTypeClient1RTT, + KeyPhase: a.keyPhase, + }) + a.qlogger.RecordEvent(qlog.KeyUpdated{ + Trigger: qlog.KeyUpdateLocal, + KeyType: qlog.KeyTypeServer1RTT, + KeyPhase: a.keyPhase, + }) + } + } + return a.keyPhase.Bit() +} + +func (a *updatableAEAD) Overhead() int { + return a.aeadOverhead +} + +func (a *updatableAEAD) EncryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) { + a.headerEncrypter.EncryptHeader(sample, firstByte, hdrBytes) +} + +func (a *updatableAEAD) DecryptHeader(sample []byte, firstByte *byte, hdrBytes []byte) { + a.headerDecrypter.DecryptHeader(sample, firstByte, hdrBytes) +} + +func (a *updatableAEAD) FirstPacketNumber() protocol.PacketNumber { + return a.firstPacketNumber +} diff --git a/third_party/quic-go/internal/handshake/updatable_aead_test.go b/third_party/quic-go/internal/handshake/updatable_aead_test.go new file mode 100644 index 0000000..84a50e3 --- /dev/null +++ b/third_party/quic-go/internal/handshake/updatable_aead_test.go @@ -0,0 +1,736 @@ +package handshake + +import ( + "crypto/fips140" + "crypto/rand" + "crypto/tls" + "fmt" + mrand "math/rand/v2" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +const ( + msg = "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua." + ad = "Donec in velit neque." +) + +func randomCipherSuite() cipherSuite { return cipherSuites[mrand.IntN(len(cipherSuites))] } + +func setupEndpoints(t *testing.T, serverRTTStats *utils.RTTStats) (client, server *updatableAEAD, serverEventRecorder *events.Recorder) { + cs := randomCipherSuite() + var eventRecorder events.Recorder + + trafficSecret1 := make([]byte, 16) + trafficSecret2 := make([]byte, 16) + rand.Read(trafficSecret1) + rand.Read(trafficSecret2) + + client = newUpdatableAEAD(utils.NewRTTStats(), nil, utils.DefaultLogger, protocol.Version1) + server = newUpdatableAEAD(serverRTTStats, &eventRecorder, utils.DefaultLogger, protocol.Version1) + client.SetReadKey(cs, trafficSecret2) + client.SetWriteKey(cs, trafficSecret1) + server.SetReadKey(cs, trafficSecret1) + server.SetWriteKey(cs, trafficSecret2) + return client, server, &eventRecorder +} + +func bothSides(ev qlogwriter.Event) []qlogwriter.Event { + switch ev := ev.(type) { + case qlog.KeyDiscarded: + return []qlogwriter.Event{ + qlog.KeyDiscarded{ + KeyType: qlog.KeyTypeClient1RTT, + KeyPhase: ev.KeyPhase, + }, + qlog.KeyDiscarded{ + KeyType: qlog.KeyTypeServer1RTT, + KeyPhase: ev.KeyPhase, + }, + } + case qlog.KeyUpdated: + return []qlogwriter.Event{ + qlog.KeyUpdated{ + KeyType: qlog.KeyTypeClient1RTT, + KeyPhase: ev.KeyPhase, + Trigger: ev.Trigger, + }, + qlog.KeyUpdated{ + KeyType: qlog.KeyTypeServer1RTT, + KeyPhase: ev.KeyPhase, + Trigger: ev.Trigger, + }, + } + default: + panic("unexpected event type: " + ev.Name()) + } +} + +func TestChaChaTestVector(t *testing.T) { + if fips140.Enabled() { + t.Skip("ChaCha20-Poly1305 is not allowed in FIPS 140-3 mode") + } + + testCases := []struct { + name string + version protocol.Version + expectedPayload []byte + expectedPacket []byte + }{ + { + version: protocol.Version1, + expectedPayload: splitHexString(t, "655e5cd55c41f69080575d7999c25a5bfb"), + expectedPacket: splitHexString(t, "4cfe4189655e5cd55c41f69080575d7999c25a5bfb"), + }, + { + version: protocol.Version2, + expectedPayload: splitHexString(t, "0ae7b6b932bc27d786f4bc2bb20f2162ba"), + expectedPacket: splitHexString(t, "5558b1c60ae7b6b932bc27d786f4bc2bb20f2162ba"), + }, + } + + for _, tc := range testCases { + t.Run(fmt.Sprintf("QUIC %s", tc.version), func(t *testing.T) { + secret := splitHexString(t, "9ac312a7f877468ebe69422748ad00a1 5443f18203a07d6060f688f30f21632b") + aead := newUpdatableAEAD(utils.NewRTTStats(), nil, nil, tc.version) + chacha := getCipherSuite(tls.TLS_CHACHA20_POLY1305_SHA256) + require.Equal(t, tls.TLS_CHACHA20_POLY1305_SHA256, chacha.ID) + aead.SetWriteKey(chacha, secret) + const pnOffset = 1 + header := splitHexString(t, "4200bff4") + payloadOffset := len(header) + plaintext := splitHexString(t, "01") + payload := aead.Seal(nil, plaintext, 654360564, header) + require.Equal(t, tc.expectedPayload, payload) + packet := append(header, payload...) + aead.EncryptHeader(packet[pnOffset+4:pnOffset+4+16], &packet[0], packet[pnOffset:payloadOffset]) + require.Equal(t, tc.expectedPacket, packet) + }) + } +} + +func TestUpdatableAEADHeaderProtection(t *testing.T) { + for _, v := range []protocol.Version{protocol.Version1, protocol.Version2} { + for _, cs := range cipherSuites { + t.Run(fmt.Sprintf("QUIC %s/%s", v, tls.CipherSuiteName(cs.ID)), func(t *testing.T) { + trafficSecret1 := make([]byte, 16) + trafficSecret2 := make([]byte, 16) + rand.Read(trafficSecret1) + rand.Read(trafficSecret2) + + client := newUpdatableAEAD(utils.NewRTTStats(), nil, utils.DefaultLogger, v) + server := newUpdatableAEAD(utils.NewRTTStats(), nil, utils.DefaultLogger, v) + client.SetReadKey(cs, trafficSecret2) + client.SetWriteKey(cs, trafficSecret1) + server.SetReadKey(cs, trafficSecret1) + server.SetWriteKey(cs, trafficSecret2) + + var lastFiveBitsDifferent int + for range 100 { + sample := make([]byte, 16) + rand.Read(sample) + header := []byte{0xb5, 1, 2, 3, 4, 5, 6, 7, 8, 0xde, 0xad, 0xbe, 0xef} + client.EncryptHeader(sample, &header[0], header[9:13]) + if header[0]&0x1f != 0xb5&0x1f { + lastFiveBitsDifferent++ + } + require.Equal(t, byte(0xb5&0xe0), header[0]&0xe0) + require.Equal(t, []byte{1, 2, 3, 4, 5, 6, 7, 8}, header[1:9]) + require.NotEqual(t, []byte{0xde, 0xad, 0xbe, 0xef}, header[9:13]) + server.DecryptHeader(sample, &header[0], header[9:13]) + require.Equal(t, []byte{0xb5, 1, 2, 3, 4, 5, 6, 7, 8, 0xde, 0xad, 0xbe, 0xef}, header) + } + require.Greater(t, lastFiveBitsDifferent, 75) + }) + } + } +} + +func TestUpdatableAEADEncryptDecryptMessage(t *testing.T) { + for _, v := range []protocol.Version{protocol.Version1, protocol.Version2} { + for _, cs := range cipherSuites { + t.Run(fmt.Sprintf("QUIC %s/%s", v, tls.CipherSuiteName(cs.ID)), func(t *testing.T) { + rttStats := utils.RTTStats{} + trafficSecret1 := make([]byte, 16) + trafficSecret2 := make([]byte, 16) + rand.Read(trafficSecret1) + rand.Read(trafficSecret2) + + client := newUpdatableAEAD(&rttStats, nil, utils.DefaultLogger, v) + server := newUpdatableAEAD(&rttStats, nil, utils.DefaultLogger, v) + client.SetReadKey(cs, trafficSecret2) + client.SetWriteKey(cs, trafficSecret1) + server.SetReadKey(cs, trafficSecret1) + server.SetWriteKey(cs, trafficSecret2) + + msg := []byte("Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.") + ad := []byte("Donec in velit neque.") + + encrypted := server.Seal(nil, msg, 0x1337, ad) + + opened, err := client.Open(nil, encrypted, monotime.Now(), 0x1337, protocol.KeyPhaseZero, ad) + require.NoError(t, err) + require.Equal(t, msg, opened) + + _, err = client.Open(nil, encrypted, monotime.Now(), 0x1337, protocol.KeyPhaseZero, []byte("wrong ad")) + require.Equal(t, ErrDecryptionFailed, err) + + _, err = client.Open(nil, encrypted, monotime.Now(), 0x42, protocol.KeyPhaseZero, ad) + require.Equal(t, ErrDecryptionFailed, err) + }) + } + } +} + +func TestUpdatableAEADPacketNumbers(t *testing.T) { + client, server, _ := setupEndpoints(t, utils.NewRTTStats()) + msg := []byte("Lorem ipsum") + ad := []byte("Donec in velit neque.") + + encrypted := server.Seal(nil, msg, 0x1337, ad) + require.Equal(t, protocol.PacketNumber(0x1337), server.FirstPacketNumber()) // make sure we save the first packet number + _ = server.Seal(nil, msg, 0x1338, ad) + require.Equal(t, protocol.PacketNumber(0x1337), server.FirstPacketNumber()) // make sure we save the first packet number + + // check that decoding the packet number works as expected + _, err := client.Open(nil, encrypted[:len(encrypted)-1], monotime.Now(), 0x1337, protocol.KeyPhaseZero, ad) + require.Error(t, err) + require.Equal(t, protocol.PacketNumber(0x38), client.DecodePacketNumber(0x38, protocol.PacketNumberLen1)) + + _, err = client.Open(nil, encrypted, monotime.Now(), 0x1337, protocol.KeyPhaseZero, ad) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumber(0x1338), client.DecodePacketNumber(0x38, protocol.PacketNumberLen1)) +} + +func TestAEADLimitReached(t *testing.T) { + client, _, _ := setupEndpoints(t, utils.NewRTTStats()) + client.invalidPacketLimit = 10 + for i := range 9 { + _, err := client.Open(nil, []byte("foobar"), monotime.Now(), protocol.PacketNumber(i), protocol.KeyPhaseZero, []byte("ad")) + require.Equal(t, ErrDecryptionFailed, err) + } + _, err := client.Open(nil, []byte("foobar"), monotime.Now(), 10, protocol.KeyPhaseZero, []byte("ad")) + require.Error(t, err) + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.AEADLimitReached, transportErr.ErrorCode) +} + +func TestKeyUpdates(t *testing.T) { + client, server, _ := setupEndpoints(t, utils.NewRTTStats()) + + now := monotime.Now() + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + encrypted0 := server.Seal(nil, []byte(msg), 0x1337, []byte(ad)) + server.rollKeys() + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + encrypted1 := server.Seal(nil, []byte(msg), 0x1337, []byte(ad)) + require.NotEqual(t, encrypted0, encrypted1) + + _, err := client.Open(nil, encrypted1, now, 0x1337, protocol.KeyPhaseZero, []byte(ad)) + require.Equal(t, ErrDecryptionFailed, err) + + client.rollKeys() + decrypted, err := client.Open(nil, encrypted1, now, 0x1337, protocol.KeyPhaseOne, []byte(ad)) + require.NoError(t, err) + require.Equal(t, msg, string(decrypted)) +} + +// func TestUpdatesKeysWhenReceivingPacketWithNextKeyPhase(t *testing.T) { +// rttStats := utils.RTTStats{} +// mockCtrl := gomock.NewController(t) +// serverTracer := mocklogging.NewMockConnectionTracer(mockCtrl) + +// trafficSecret1 := make([]byte, 16) +// trafficSecret2 := make([]byte, 16) +// rand.Read(trafficSecret1) +// rand.Read(trafficSecret2) + +// client := newUpdatableAEAD(&rttStats, nil, utils.DefaultLogger, protocol.Version1) +// server := newUpdatableAEAD(&rttStats, serverTracer, utils.DefaultLogger, protocol.Version1) +// client.SetReadKey(cs, trafficSecret2) +// client.SetWriteKey(cs, trafficSecret1) +// server.SetReadKey(cs, trafficSecret1) +// server.SetWriteKey(cs, trafficSecret2) + +// now := monotime.Now() +// encrypted0 := client.Seal(nil, []byte(msg), 0x42, ad) +// decrypted, err := server.Open(nil, encrypted0, now, 0x42, protocol.KeyPhaseZero, ad) +// require.NoError(t, err) +// require.Equal(t, msg, decrypted) + +// require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) +// _ = server.Seal(nil, msg, 0x1, ad) + +// client.rollKeys() +// encrypted1 := client.Seal(nil, msg, 0x43, ad) +// serverTracer.EXPECT().UpdatedKey(protocol.KeyPhase(1), true) +// decrypted, err = server.Open(nil, encrypted1, now, 0x43, protocol.KeyPhaseOne, ad) +// require.NoError(t, err) +// require.Equal(t, msg, decrypted) +// require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) +// } + +func TestReorderedPacketAfterKeyUpdate(t *testing.T) { + client, server, eventRecorder := setupEndpoints(t, utils.NewRTTStats()) + + now := monotime.Now() + encrypted01 := client.Seal(nil, []byte(msg), 0x42, []byte(ad)) + encrypted02 := client.Seal(nil, []byte(msg), 0x43, []byte(ad)) + _, err := server.Open(nil, encrypted01, now, 0x42, protocol.KeyPhaseZero, []byte(ad)) + require.NoError(t, err) + _ = server.Seal(nil, []byte(msg), 0x1, []byte(ad)) + + client.rollKeys() + encrypted1 := client.Seal(nil, []byte(msg), 0x44, []byte(ad)) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + _, err = server.Open(nil, encrypted1, now, 0x44, protocol.KeyPhaseOne, []byte(ad)) + require.NoError(t, err) + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{Trigger: qlog.KeyUpdateRemote, KeyPhase: 1}), + eventRecorder.Events(), + ) + + // now receive a reordered packet + decrypted, err := server.Open(nil, encrypted02, now, 0x43, protocol.KeyPhaseZero, []byte(ad)) + require.NoError(t, err) + require.Equal(t, msg, string(decrypted)) + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) +} + +func TestDropsKeys3PTOsAfterKeyUpdate(t *testing.T) { + rttStats := utils.NewRTTStats() + client, server, eventRecorder := setupEndpoints(t, rttStats) + + now := monotime.Now() + rttStats.UpdateRTT(10*time.Millisecond, 0) + pto := rttStats.PTO(true) + encrypted01 := client.Seal(nil, []byte(msg), 0x42, []byte(ad)) + encrypted02 := client.Seal(nil, []byte(msg), 0x43, []byte(ad)) + _, err := server.Open(nil, encrypted01, now, 0x42, protocol.KeyPhaseZero, []byte(ad)) + require.NoError(t, err) + _ = server.Seal(nil, []byte(msg), 0x1, []byte(ad)) + + client.rollKeys() + encrypted1 := client.Seal(nil, []byte(msg), 0x44, []byte(ad)) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + _, err = server.Open(nil, encrypted1, now, 0x44, protocol.KeyPhaseOne, []byte(ad)) + require.NoError(t, err) + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateRemote}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // packet arrived too late, the key was already dropped + _, err = server.Open(nil, encrypted02, now.Add(3*pto).Add(time.Nanosecond), 0x43, protocol.KeyPhaseZero, []byte(ad)) + require.Equal(t, ErrKeysDropped, err) + require.Equal(t, + bothSides(qlog.KeyDiscarded{KeyPhase: 0}), + eventRecorder.Events(), + ) +} + +func TestAllowsFirstKeyUpdateImmediately(t *testing.T) { + client, server, serverTracer := setupEndpoints(t, utils.NewRTTStats()) + client.rollKeys() + encrypted := client.Seal(nil, []byte(msg), 0x1337, []byte(ad)) + + // if decryption failed, we don't expect a key phase update + _, err := server.Open(nil, encrypted[:len(encrypted)-1], monotime.Now(), 0x1337, protocol.KeyPhaseOne, []byte(ad)) + require.Equal(t, ErrDecryptionFailed, err) + + // the key phase is updated on first successful decryption + _, err = server.Open(nil, encrypted, monotime.Now(), 0x1337, protocol.KeyPhaseOne, []byte(ad)) + require.NoError(t, err) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateRemote}), + serverTracer.Events(), + ) +} + +func TestRejectFrequentKeyUpdates(t *testing.T) { + client, server, _ := setupEndpoints(t, utils.NewRTTStats()) + + server.rollKeys() + client.rollKeys() + encrypted0 := client.Seal(nil, []byte(msg), 0x42, []byte(ad)) + _, err := server.Open(nil, encrypted0, monotime.Now(), 0x42, protocol.KeyPhaseOne, []byte(ad)) + require.NoError(t, err) + + client.rollKeys() + encrypted1 := client.Seal(nil, []byte(msg), 0x42, []byte(ad)) + _, err = server.Open(nil, encrypted1, monotime.Now(), 0x42, protocol.KeyPhaseZero, []byte(ad)) + require.Equal(t, &qerr.TransportError{ + ErrorCode: qerr.KeyUpdateError, + ErrorMessage: "keys updated too quickly", + }, err) +} + +func setKeyUpdateIntervals(t *testing.T, firstKeyUpdateInterval, keyUpdateInterval uint64) { + reset := SetKeyUpdateInterval(keyUpdateInterval) + t.Cleanup(reset) + + origFirstKeyUpdateInterval := FirstKeyUpdateInterval + FirstKeyUpdateInterval = firstKeyUpdateInterval + + t.Cleanup(func() { FirstKeyUpdateInterval = origFirstKeyUpdateInterval }) +} + +func TestInitiateKeyUpdateAfterSendingMaxPackets(t *testing.T) { + const firstKeyUpdateInterval = 5 + const keyUpdateInterval = 20 + setKeyUpdateIntervals(t, firstKeyUpdateInterval, keyUpdateInterval) + + client, server, eventRecorder := setupEndpoints(t, utils.NewRTTStats()) + server.SetHandshakeConfirmed() + + var pn protocol.PacketNumber + // first key update + for range firstKeyUpdateInterval { + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + pn++ + } + // the first update is allowed without receiving an acknowledgement + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateLocal}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // subsequent key update + for range 2 * keyUpdateInterval { + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + pn++ + } + // no update allowed before receiving an acknowledgement for the current key phase + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + // receive an ACK for a packet sent in key phase 1 + client.rollKeys() + b := client.Seal(nil, []byte("foobar"), 1, []byte("ad")) + _, err := server.Open(nil, b, monotime.Now(), 1, protocol.KeyPhaseOne, []byte("ad")) + require.NoError(t, err) + require.NoError(t, server.SetLargestAcked(firstKeyUpdateInterval)) + require.Empty(t, eventRecorder.Events()) + + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + require.Equal(t, + append( + bothSides(qlog.KeyDiscarded{KeyPhase: 0}), + bothSides(qlog.KeyUpdated{KeyPhase: 2, Trigger: qlog.KeyUpdateLocal})..., + ), + eventRecorder.Events(), + ) +} + +func TestKeyUpdateEnforceACKKeyPhase(t *testing.T) { + const firstKeyUpdateInterval = 5 + setKeyUpdateIntervals(t, firstKeyUpdateInterval, protocol.KeyUpdateInterval) + + _, server, eventRecorder := setupEndpoints(t, utils.NewRTTStats()) + server.SetHandshakeConfirmed() + + // First make sure that we update our keys. + for i := range firstKeyUpdateInterval { + pn := protocol.PacketNumber(i) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + } + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateLocal}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // Now that our keys are updated, send a packet using the new keys. + const nextPN = firstKeyUpdateInterval + 1 + server.Seal(nil, []byte(msg), nextPN, []byte(ad)) + + for i := range firstKeyUpdateInterval { + // We haven't decrypted any packet in the new key phase yet. + // This means that the ACK must have been sent in the old key phase. + require.NoError(t, server.SetLargestAcked(protocol.PacketNumber(i))) + } + + // We haven't decrypted any packet in the new key phase yet. + // This means that the ACK must have been sent in the old key phase. + err := server.SetLargestAcked(nextPN) + require.Error(t, err) + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.KeyUpdateError, transportErr.ErrorCode) + require.Equal(t, "received ACK for key phase 1, but peer didn't update keys", transportErr.ErrorMessage) + require.Empty(t, eventRecorder.Events()) +} + +func TestKeyUpdateAfterOpeningMaxPackets(t *testing.T) { + const firstKeyUpdateInterval = 5 + const keyUpdateInterval = 20 + setKeyUpdateIntervals(t, firstKeyUpdateInterval, keyUpdateInterval) + + client, server, eventRecorder := setupEndpoints(t, utils.NewRTTStats()) + server.SetHandshakeConfirmed() + + msg := []byte("message") + ad := []byte("additional data") + + // first key update + var pn protocol.PacketNumber + for range firstKeyUpdateInterval { + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + encrypted := client.Seal(nil, msg, pn, ad) + _, err := server.Open(nil, encrypted, monotime.Now(), pn, protocol.KeyPhaseZero, ad) + require.NoError(t, err) + pn++ + } + + // the first update is allowed without receiving an acknowledgement + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateLocal}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // subsequent key update + client.rollKeys() + for range keyUpdateInterval { + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + encrypted := client.Seal(nil, msg, pn, ad) + _, err := server.Open(nil, encrypted, monotime.Now(), pn, protocol.KeyPhaseOne, ad) + require.NoError(t, err) + pn++ + } + + // No update allowed before receiving an acknowledgement for the current key phase + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + server.Seal(nil, msg, 1, ad) + require.NoError(t, server.SetLargestAcked(firstKeyUpdateInterval+1)) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + require.Equal(t, + append( + bothSides(qlog.KeyDiscarded{KeyPhase: 0}), + bothSides(qlog.KeyUpdated{KeyPhase: 2, Trigger: qlog.KeyUpdateLocal})..., + ), + eventRecorder.Events(), + ) +} + +func TestKeyUpdateKeyPhaseSkipping(t *testing.T) { + const firstKeyUpdateInterval = 5 + const keyUpdateInterval = 20 + setKeyUpdateIntervals(t, firstKeyUpdateInterval, keyUpdateInterval) + + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(10*time.Millisecond, 0) + client, server, eventRecorder := setupEndpoints(t, rttStats) + server.SetHandshakeConfirmed() + + now := monotime.Now() + data1 := client.Seal(nil, []byte(msg), 1, []byte(ad)) + _, err := server.Open(nil, data1, now, 1, protocol.KeyPhaseZero, []byte(ad)) + require.NoError(t, err) + for i := range firstKeyUpdateInterval { + pn := protocol.PacketNumber(i) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + require.NoError(t, server.SetLargestAcked(pn)) + } + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateLocal}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // The server never received a packet at key phase 1. + // Make sure the key phase 0 is still there at a much later point. + data2 := client.Seal(nil, []byte(msg), 2, []byte(ad)) + _, err = server.Open(nil, data2, now.Add(10*rttStats.PTO(true)), 2, protocol.KeyPhaseZero, []byte(ad)) + require.NoError(t, err) + require.Empty(t, eventRecorder.Events()) +} + +func TestFastKeyUpdatesByPeer(t *testing.T) { + const firstKeyUpdateInterval = 5 + const keyUpdateInterval = 20 + setKeyUpdateIntervals(t, firstKeyUpdateInterval, keyUpdateInterval) + + client, server, eventRecorder := setupEndpoints(t, utils.NewRTTStats()) + server.SetHandshakeConfirmed() + + var pn protocol.PacketNumber + for range firstKeyUpdateInterval { + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + pn++ + } + b := client.Seal(nil, []byte("foobar"), 1, []byte("ad")) + _, err := server.Open(nil, b, monotime.Now(), 1, protocol.KeyPhaseZero, []byte("ad")) + require.NoError(t, err) + require.NoError(t, server.SetLargestAcked(0)) + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateLocal}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // Send and receive an acknowledgement for a packet in key phase 1. + // We are now running a timer to drop the keys with 3 PTO. + server.Seal(nil, []byte(msg), pn, []byte(ad)) + client.rollKeys() + dataKeyPhaseOne := client.Seal(nil, []byte(msg), 2, []byte(ad)) + now := monotime.Now() + _, err = server.Open(nil, dataKeyPhaseOne, now, 2, protocol.KeyPhaseOne, []byte(ad)) + require.NoError(t, err) + require.NoError(t, server.SetLargestAcked(pn)) + // Now the client sends us a packet in key phase 2, forcing us to update keys before the 3 PTO period is over. + // This mean that we need to drop the keys for key phase 0 immediately. + client.rollKeys() + dataKeyPhaseTwo := client.Seal(nil, []byte(msg), 3, []byte(ad)) + + _, err = server.Open(nil, dataKeyPhaseTwo, now, 3, protocol.KeyPhaseZero, []byte(ad)) + require.NoError(t, err) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + require.Equal(t, + append( + bothSides(qlog.KeyDiscarded{KeyPhase: 0}), + bothSides(qlog.KeyUpdated{KeyPhase: 2, Trigger: qlog.KeyUpdateRemote})..., + ), + eventRecorder.Events(), + ) +} + +func TestFastKeyUpdateByUs(t *testing.T) { + const firstKeyUpdateInterval = 5 + const keyUpdateInterval = 20 + setKeyUpdateIntervals(t, firstKeyUpdateInterval, keyUpdateInterval) + + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(10*time.Millisecond, 0) + client, server, eventRecorder := setupEndpoints(t, rttStats) + server.SetHandshakeConfirmed() + + // send so many packets that we initiate the first key update + for i := range firstKeyUpdateInterval { + pn := protocol.PacketNumber(i) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + } + b := client.Seal(nil, []byte("foobar"), 1, []byte("ad")) + _, err := server.Open(nil, b, monotime.Now(), 1, protocol.KeyPhaseZero, []byte("ad")) + require.NoError(t, err) + require.NoError(t, server.SetLargestAcked(0)) + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + require.Equal(t, + bothSides(qlog.KeyUpdated{KeyPhase: 1, Trigger: qlog.KeyUpdateLocal}), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // send so many packets that we initiate the next key update + for i := keyUpdateInterval; i < 2*keyUpdateInterval; i++ { + pn := protocol.PacketNumber(i) + require.Equal(t, protocol.KeyPhaseOne, server.KeyPhase()) + server.Seal(nil, []byte(msg), pn, []byte(ad)) + } + client.rollKeys() + b = client.Seal(nil, []byte("foobar"), 2, []byte("ad")) + now := monotime.Now() + _, err = server.Open(nil, b, now, 2, protocol.KeyPhaseOne, []byte("ad")) + require.NoError(t, err) + require.NoError(t, server.SetLargestAcked(keyUpdateInterval)) + require.Equal(t, protocol.KeyPhaseZero, server.KeyPhase()) + require.Equal(t, + append( + bothSides(qlog.KeyDiscarded{KeyPhase: 0}), + bothSides(qlog.KeyUpdated{KeyPhase: 2, Trigger: qlog.KeyUpdateLocal})..., + ), + eventRecorder.Events(), + ) + eventRecorder.Clear() + + // We haven't received an ACK for a packet sent in key phase 2 yet. + // Make sure we canceled the timer to drop the previous key phase. + b = client.Seal(nil, []byte("foobar"), 3, []byte("ad")) + _, err = server.Open(nil, b, now.Add(10*rttStats.PTO(true)), 3, protocol.KeyPhaseOne, []byte("ad")) + require.NoError(t, err) + require.Empty(t, eventRecorder.Events()) +} + +func getClientAndServer() (client, server *updatableAEAD) { + trafficSecret1 := make([]byte, 16) + trafficSecret2 := make([]byte, 16) + rand.Read(trafficSecret1) + rand.Read(trafficSecret2) + + cs := cipherSuites[0] + rttStats := utils.NewRTTStats() + client = newUpdatableAEAD(rttStats, nil, utils.DefaultLogger, protocol.Version1) + server = newUpdatableAEAD(rttStats, nil, utils.DefaultLogger, protocol.Version1) + client.SetReadKey(cs, trafficSecret2) + client.SetWriteKey(cs, trafficSecret1) + server.SetReadKey(cs, trafficSecret1) + server.SetWriteKey(cs, trafficSecret2) + return +} + +func BenchmarkPacketEncryption(b *testing.B) { + client, _ := getClientAndServer() + const l = 1200 + src := make([]byte, l) + rand.Read(src) + ad := make([]byte, 32) + rand.Read(ad) + + var pn protocol.PacketNumber + for b.Loop() { + src = client.Seal(src[:0], src[:l], pn, ad) + pn++ + } +} + +func BenchmarkPacketDecryption(b *testing.B) { + client, server := getClientAndServer() + const l = 1200 + src := make([]byte, l) + dst := make([]byte, l) + rand.Read(src) + ad := make([]byte, 32) + rand.Read(ad) + src = client.Seal(src[:0], src[:l], 1337, ad) + + for b.Loop() { + if _, err := server.Open(dst[:0], src, 0, 1337, protocol.KeyPhaseZero, ad); err != nil { + b.Fatalf("opening failed: %v", err) + } + } +} + +func BenchmarkRollKeys(b *testing.B) { + client, _ := getClientAndServer() + + for b.Loop() { + client.rollKeys() + } + if int(client.keyPhase) != b.N { + b.Fatal("didn't roll keys often enough") + } +} diff --git a/third_party/quic-go/internal/mocks/ackhandler/sent_packet_handler.go b/third_party/quic-go/internal/mocks/ackhandler/sent_packet_handler.go new file mode 100644 index 0000000..990010a --- /dev/null +++ b/third_party/quic-go/internal/mocks/ackhandler/sent_packet_handler.go @@ -0,0 +1,713 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/ackhandler (interfaces: SentPacketHandler) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package mockackhandler -destination ackhandler/sent_packet_handler.go github.com/apernet/quic-go/internal/ackhandler SentPacketHandler +// + +// Package mockackhandler is a generated GoMock package. +package mockackhandler + +import ( + reflect "reflect" + + congestion "github.com/apernet/quic-go/congestion" + ackhandler "github.com/apernet/quic-go/internal/ackhandler" + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + wire "github.com/apernet/quic-go/internal/wire" + gomock "go.uber.org/mock/gomock" +) + +// MockSentPacketHandler is a mock of SentPacketHandler interface. +type MockSentPacketHandler struct { + ctrl *gomock.Controller + recorder *MockSentPacketHandlerMockRecorder + isgomock struct{} +} + +// MockSentPacketHandlerMockRecorder is the mock recorder for MockSentPacketHandler. +type MockSentPacketHandlerMockRecorder struct { + mock *MockSentPacketHandler +} + +// NewMockSentPacketHandler creates a new mock instance. +func NewMockSentPacketHandler(ctrl *gomock.Controller) *MockSentPacketHandler { + mock := &MockSentPacketHandler{ctrl: ctrl} + mock.recorder = &MockSentPacketHandlerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockSentPacketHandler) EXPECT() *MockSentPacketHandlerMockRecorder { + return m.recorder +} + +// DropPackets mocks base method. +func (m *MockSentPacketHandler) DropPackets(arg0 protocol.EncryptionLevel, rcvTime monotime.Time) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "DropPackets", arg0, rcvTime) +} + +// DropPackets indicates an expected call of DropPackets. +func (mr *MockSentPacketHandlerMockRecorder) DropPackets(arg0, rcvTime any) *MockSentPacketHandlerDropPacketsCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DropPackets", reflect.TypeOf((*MockSentPacketHandler)(nil).DropPackets), arg0, rcvTime) + return &MockSentPacketHandlerDropPacketsCall{Call: call} +} + +// MockSentPacketHandlerDropPacketsCall wrap *gomock.Call +type MockSentPacketHandlerDropPacketsCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerDropPacketsCall) Return() *MockSentPacketHandlerDropPacketsCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerDropPacketsCall) Do(f func(protocol.EncryptionLevel, monotime.Time)) *MockSentPacketHandlerDropPacketsCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerDropPacketsCall) DoAndReturn(f func(protocol.EncryptionLevel, monotime.Time)) *MockSentPacketHandlerDropPacketsCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ECNMode mocks base method. +func (m *MockSentPacketHandler) ECNMode(isShortHeaderPacket bool) protocol.ECN { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ECNMode", isShortHeaderPacket) + ret0, _ := ret[0].(protocol.ECN) + return ret0 +} + +// ECNMode indicates an expected call of ECNMode. +func (mr *MockSentPacketHandlerMockRecorder) ECNMode(isShortHeaderPacket any) *MockSentPacketHandlerECNModeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ECNMode", reflect.TypeOf((*MockSentPacketHandler)(nil).ECNMode), isShortHeaderPacket) + return &MockSentPacketHandlerECNModeCall{Call: call} +} + +// MockSentPacketHandlerECNModeCall wrap *gomock.Call +type MockSentPacketHandlerECNModeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerECNModeCall) Return(arg0 protocol.ECN) *MockSentPacketHandlerECNModeCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerECNModeCall) Do(f func(bool) protocol.ECN) *MockSentPacketHandlerECNModeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerECNModeCall) DoAndReturn(f func(bool) protocol.ECN) *MockSentPacketHandlerECNModeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetLossDetectionTimeout mocks base method. +func (m *MockSentPacketHandler) GetLossDetectionTimeout() monotime.Time { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetLossDetectionTimeout") + ret0, _ := ret[0].(monotime.Time) + return ret0 +} + +// GetLossDetectionTimeout indicates an expected call of GetLossDetectionTimeout. +func (mr *MockSentPacketHandlerMockRecorder) GetLossDetectionTimeout() *MockSentPacketHandlerGetLossDetectionTimeoutCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetLossDetectionTimeout", reflect.TypeOf((*MockSentPacketHandler)(nil).GetLossDetectionTimeout)) + return &MockSentPacketHandlerGetLossDetectionTimeoutCall{Call: call} +} + +// MockSentPacketHandlerGetLossDetectionTimeoutCall wrap *gomock.Call +type MockSentPacketHandlerGetLossDetectionTimeoutCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerGetLossDetectionTimeoutCall) Return(arg0 monotime.Time) *MockSentPacketHandlerGetLossDetectionTimeoutCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerGetLossDetectionTimeoutCall) Do(f func() monotime.Time) *MockSentPacketHandlerGetLossDetectionTimeoutCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerGetLossDetectionTimeoutCall) DoAndReturn(f func() monotime.Time) *MockSentPacketHandlerGetLossDetectionTimeoutCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// MigratedPath mocks base method. +func (m *MockSentPacketHandler) MigratedPath(now monotime.Time, initialMaxPacketSize protocol.ByteCount) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "MigratedPath", now, initialMaxPacketSize) +} + +// MigratedPath indicates an expected call of MigratedPath. +func (mr *MockSentPacketHandlerMockRecorder) MigratedPath(now, initialMaxPacketSize any) *MockSentPacketHandlerMigratedPathCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "MigratedPath", reflect.TypeOf((*MockSentPacketHandler)(nil).MigratedPath), now, initialMaxPacketSize) + return &MockSentPacketHandlerMigratedPathCall{Call: call} +} + +// MockSentPacketHandlerMigratedPathCall wrap *gomock.Call +type MockSentPacketHandlerMigratedPathCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerMigratedPathCall) Return() *MockSentPacketHandlerMigratedPathCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerMigratedPathCall) Do(f func(monotime.Time, protocol.ByteCount)) *MockSentPacketHandlerMigratedPathCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerMigratedPathCall) DoAndReturn(f func(monotime.Time, protocol.ByteCount)) *MockSentPacketHandlerMigratedPathCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// OnLossDetectionTimeout mocks base method. +func (m *MockSentPacketHandler) OnLossDetectionTimeout(now monotime.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "OnLossDetectionTimeout", now) + ret0, _ := ret[0].(error) + return ret0 +} + +// OnLossDetectionTimeout indicates an expected call of OnLossDetectionTimeout. +func (mr *MockSentPacketHandlerMockRecorder) OnLossDetectionTimeout(now any) *MockSentPacketHandlerOnLossDetectionTimeoutCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnLossDetectionTimeout", reflect.TypeOf((*MockSentPacketHandler)(nil).OnLossDetectionTimeout), now) + return &MockSentPacketHandlerOnLossDetectionTimeoutCall{Call: call} +} + +// MockSentPacketHandlerOnLossDetectionTimeoutCall wrap *gomock.Call +type MockSentPacketHandlerOnLossDetectionTimeoutCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerOnLossDetectionTimeoutCall) Return(arg0 error) *MockSentPacketHandlerOnLossDetectionTimeoutCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerOnLossDetectionTimeoutCall) Do(f func(monotime.Time) error) *MockSentPacketHandlerOnLossDetectionTimeoutCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerOnLossDetectionTimeoutCall) DoAndReturn(f func(monotime.Time) error) *MockSentPacketHandlerOnLossDetectionTimeoutCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PeekPacketNumber mocks base method. +func (m *MockSentPacketHandler) PeekPacketNumber(arg0 protocol.EncryptionLevel) (protocol.PacketNumber, protocol.PacketNumberLen) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PeekPacketNumber", arg0) + ret0, _ := ret[0].(protocol.PacketNumber) + ret1, _ := ret[1].(protocol.PacketNumberLen) + return ret0, ret1 +} + +// PeekPacketNumber indicates an expected call of PeekPacketNumber. +func (mr *MockSentPacketHandlerMockRecorder) PeekPacketNumber(arg0 any) *MockSentPacketHandlerPeekPacketNumberCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PeekPacketNumber", reflect.TypeOf((*MockSentPacketHandler)(nil).PeekPacketNumber), arg0) + return &MockSentPacketHandlerPeekPacketNumberCall{Call: call} +} + +// MockSentPacketHandlerPeekPacketNumberCall wrap *gomock.Call +type MockSentPacketHandlerPeekPacketNumberCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerPeekPacketNumberCall) Return(arg0 protocol.PacketNumber, arg1 protocol.PacketNumberLen) *MockSentPacketHandlerPeekPacketNumberCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerPeekPacketNumberCall) Do(f func(protocol.EncryptionLevel) (protocol.PacketNumber, protocol.PacketNumberLen)) *MockSentPacketHandlerPeekPacketNumberCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerPeekPacketNumberCall) DoAndReturn(f func(protocol.EncryptionLevel) (protocol.PacketNumber, protocol.PacketNumberLen)) *MockSentPacketHandlerPeekPacketNumberCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PopPacketNumber mocks base method. +func (m *MockSentPacketHandler) PopPacketNumber(arg0 protocol.EncryptionLevel) protocol.PacketNumber { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PopPacketNumber", arg0) + ret0, _ := ret[0].(protocol.PacketNumber) + return ret0 +} + +// PopPacketNumber indicates an expected call of PopPacketNumber. +func (mr *MockSentPacketHandlerMockRecorder) PopPacketNumber(arg0 any) *MockSentPacketHandlerPopPacketNumberCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PopPacketNumber", reflect.TypeOf((*MockSentPacketHandler)(nil).PopPacketNumber), arg0) + return &MockSentPacketHandlerPopPacketNumberCall{Call: call} +} + +// MockSentPacketHandlerPopPacketNumberCall wrap *gomock.Call +type MockSentPacketHandlerPopPacketNumberCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerPopPacketNumberCall) Return(arg0 protocol.PacketNumber) *MockSentPacketHandlerPopPacketNumberCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerPopPacketNumberCall) Do(f func(protocol.EncryptionLevel) protocol.PacketNumber) *MockSentPacketHandlerPopPacketNumberCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerPopPacketNumberCall) DoAndReturn(f func(protocol.EncryptionLevel) protocol.PacketNumber) *MockSentPacketHandlerPopPacketNumberCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// QueueProbePacket mocks base method. +func (m *MockSentPacketHandler) QueueProbePacket(arg0 protocol.EncryptionLevel) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "QueueProbePacket", arg0) + ret0, _ := ret[0].(bool) + return ret0 +} + +// QueueProbePacket indicates an expected call of QueueProbePacket. +func (mr *MockSentPacketHandlerMockRecorder) QueueProbePacket(arg0 any) *MockSentPacketHandlerQueueProbePacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "QueueProbePacket", reflect.TypeOf((*MockSentPacketHandler)(nil).QueueProbePacket), arg0) + return &MockSentPacketHandlerQueueProbePacketCall{Call: call} +} + +// MockSentPacketHandlerQueueProbePacketCall wrap *gomock.Call +type MockSentPacketHandlerQueueProbePacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerQueueProbePacketCall) Return(arg0 bool) *MockSentPacketHandlerQueueProbePacketCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerQueueProbePacketCall) Do(f func(protocol.EncryptionLevel) bool) *MockSentPacketHandlerQueueProbePacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerQueueProbePacketCall) DoAndReturn(f func(protocol.EncryptionLevel) bool) *MockSentPacketHandlerQueueProbePacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReceivedAck mocks base method. +func (m *MockSentPacketHandler) ReceivedAck(f *wire.AckFrame, encLevel protocol.EncryptionLevel, rcvTime monotime.Time) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReceivedAck", f, encLevel, rcvTime) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ReceivedAck indicates an expected call of ReceivedAck. +func (mr *MockSentPacketHandlerMockRecorder) ReceivedAck(f, encLevel, rcvTime any) *MockSentPacketHandlerReceivedAckCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReceivedAck", reflect.TypeOf((*MockSentPacketHandler)(nil).ReceivedAck), f, encLevel, rcvTime) + return &MockSentPacketHandlerReceivedAckCall{Call: call} +} + +// MockSentPacketHandlerReceivedAckCall wrap *gomock.Call +type MockSentPacketHandlerReceivedAckCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerReceivedAckCall) Return(arg0 bool, arg1 error) *MockSentPacketHandlerReceivedAckCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerReceivedAckCall) Do(f func(*wire.AckFrame, protocol.EncryptionLevel, monotime.Time) (bool, error)) *MockSentPacketHandlerReceivedAckCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerReceivedAckCall) DoAndReturn(f func(*wire.AckFrame, protocol.EncryptionLevel, monotime.Time) (bool, error)) *MockSentPacketHandlerReceivedAckCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReceivedBytes mocks base method. +func (m *MockSentPacketHandler) ReceivedBytes(arg0 protocol.ByteCount, rcvTime monotime.Time) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "ReceivedBytes", arg0, rcvTime) +} + +// ReceivedBytes indicates an expected call of ReceivedBytes. +func (mr *MockSentPacketHandlerMockRecorder) ReceivedBytes(arg0, rcvTime any) *MockSentPacketHandlerReceivedBytesCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReceivedBytes", reflect.TypeOf((*MockSentPacketHandler)(nil).ReceivedBytes), arg0, rcvTime) + return &MockSentPacketHandlerReceivedBytesCall{Call: call} +} + +// MockSentPacketHandlerReceivedBytesCall wrap *gomock.Call +type MockSentPacketHandlerReceivedBytesCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerReceivedBytesCall) Return() *MockSentPacketHandlerReceivedBytesCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerReceivedBytesCall) Do(f func(protocol.ByteCount, monotime.Time)) *MockSentPacketHandlerReceivedBytesCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerReceivedBytesCall) DoAndReturn(f func(protocol.ByteCount, monotime.Time)) *MockSentPacketHandlerReceivedBytesCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReceivedPacket mocks base method. +func (m *MockSentPacketHandler) ReceivedPacket(arg0 protocol.EncryptionLevel, arg1 monotime.Time) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "ReceivedPacket", arg0, arg1) +} + +// ReceivedPacket indicates an expected call of ReceivedPacket. +func (mr *MockSentPacketHandlerMockRecorder) ReceivedPacket(arg0, arg1 any) *MockSentPacketHandlerReceivedPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReceivedPacket", reflect.TypeOf((*MockSentPacketHandler)(nil).ReceivedPacket), arg0, arg1) + return &MockSentPacketHandlerReceivedPacketCall{Call: call} +} + +// MockSentPacketHandlerReceivedPacketCall wrap *gomock.Call +type MockSentPacketHandlerReceivedPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerReceivedPacketCall) Return() *MockSentPacketHandlerReceivedPacketCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerReceivedPacketCall) Do(f func(protocol.EncryptionLevel, monotime.Time)) *MockSentPacketHandlerReceivedPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerReceivedPacketCall) DoAndReturn(f func(protocol.EncryptionLevel, monotime.Time)) *MockSentPacketHandlerReceivedPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ResetForRetry mocks base method. +func (m *MockSentPacketHandler) ResetForRetry(rcvTime monotime.Time) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "ResetForRetry", rcvTime) +} + +// ResetForRetry indicates an expected call of ResetForRetry. +func (mr *MockSentPacketHandlerMockRecorder) ResetForRetry(rcvTime any) *MockSentPacketHandlerResetForRetryCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ResetForRetry", reflect.TypeOf((*MockSentPacketHandler)(nil).ResetForRetry), rcvTime) + return &MockSentPacketHandlerResetForRetryCall{Call: call} +} + +// MockSentPacketHandlerResetForRetryCall wrap *gomock.Call +type MockSentPacketHandlerResetForRetryCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerResetForRetryCall) Return() *MockSentPacketHandlerResetForRetryCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerResetForRetryCall) Do(f func(monotime.Time)) *MockSentPacketHandlerResetForRetryCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerResetForRetryCall) DoAndReturn(f func(monotime.Time)) *MockSentPacketHandlerResetForRetryCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SendMode mocks base method. +func (m *MockSentPacketHandler) SendMode(now monotime.Time) ackhandler.SendMode { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SendMode", now) + ret0, _ := ret[0].(ackhandler.SendMode) + return ret0 +} + +// SendMode indicates an expected call of SendMode. +func (mr *MockSentPacketHandlerMockRecorder) SendMode(now any) *MockSentPacketHandlerSendModeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SendMode", reflect.TypeOf((*MockSentPacketHandler)(nil).SendMode), now) + return &MockSentPacketHandlerSendModeCall{Call: call} +} + +// MockSentPacketHandlerSendModeCall wrap *gomock.Call +type MockSentPacketHandlerSendModeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerSendModeCall) Return(arg0 ackhandler.SendMode) *MockSentPacketHandlerSendModeCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerSendModeCall) Do(f func(monotime.Time) ackhandler.SendMode) *MockSentPacketHandlerSendModeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerSendModeCall) DoAndReturn(f func(monotime.Time) ackhandler.SendMode) *MockSentPacketHandlerSendModeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SentPacket mocks base method. +func (m *MockSentPacketHandler) SentPacket(t monotime.Time, pn, largestAcked protocol.PacketNumber, streamFrames []ackhandler.StreamFrame, frames []ackhandler.Frame, encLevel protocol.EncryptionLevel, ecn protocol.ECN, size protocol.ByteCount, isPathMTUProbePacket, isPathProbePacket bool) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SentPacket", t, pn, largestAcked, streamFrames, frames, encLevel, ecn, size, isPathMTUProbePacket, isPathProbePacket) +} + +// SentPacket indicates an expected call of SentPacket. +func (mr *MockSentPacketHandlerMockRecorder) SentPacket(t, pn, largestAcked, streamFrames, frames, encLevel, ecn, size, isPathMTUProbePacket, isPathProbePacket any) *MockSentPacketHandlerSentPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SentPacket", reflect.TypeOf((*MockSentPacketHandler)(nil).SentPacket), t, pn, largestAcked, streamFrames, frames, encLevel, ecn, size, isPathMTUProbePacket, isPathProbePacket) + return &MockSentPacketHandlerSentPacketCall{Call: call} +} + +// MockSentPacketHandlerSentPacketCall wrap *gomock.Call +type MockSentPacketHandlerSentPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerSentPacketCall) Return() *MockSentPacketHandlerSentPacketCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerSentPacketCall) Do(f func(monotime.Time, protocol.PacketNumber, protocol.PacketNumber, []ackhandler.StreamFrame, []ackhandler.Frame, protocol.EncryptionLevel, protocol.ECN, protocol.ByteCount, bool, bool)) *MockSentPacketHandlerSentPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerSentPacketCall) DoAndReturn(f func(monotime.Time, protocol.PacketNumber, protocol.PacketNumber, []ackhandler.StreamFrame, []ackhandler.Frame, protocol.EncryptionLevel, protocol.ECN, protocol.ByteCount, bool, bool)) *MockSentPacketHandlerSentPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetCongestionControl mocks base method. +func (m *MockSentPacketHandler) SetCongestionControl(arg0 congestion.CongestionControl) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetCongestionControl", arg0) +} + +// SetCongestionControl indicates an expected call of SetCongestionControl. +func (mr *MockSentPacketHandlerMockRecorder) SetCongestionControl(arg0 any) *MockSentPacketHandlerSetCongestionControlCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetCongestionControl", reflect.TypeOf((*MockSentPacketHandler)(nil).SetCongestionControl), arg0) + return &MockSentPacketHandlerSetCongestionControlCall{Call: call} +} + +// MockSentPacketHandlerSetCongestionControlCall wrap *gomock.Call +type MockSentPacketHandlerSetCongestionControlCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerSetCongestionControlCall) Return() *MockSentPacketHandlerSetCongestionControlCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerSetCongestionControlCall) Do(f func(congestion.CongestionControl)) *MockSentPacketHandlerSetCongestionControlCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerSetCongestionControlCall) DoAndReturn(f func(congestion.CongestionControl)) *MockSentPacketHandlerSetCongestionControlCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetLastDatagramPadding mocks base method. +func (m *MockSentPacketHandler) SetLastDatagramPadding(arg0 protocol.ByteCount) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetLastDatagramPadding", arg0) +} + +// SetLastDatagramPadding indicates an expected call of SetLastDatagramPadding. +func (mr *MockSentPacketHandlerMockRecorder) SetLastDatagramPadding(arg0 any) *MockSentPacketHandlerSetLastDatagramPaddingCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetLastDatagramPadding", reflect.TypeOf((*MockSentPacketHandler)(nil).SetLastDatagramPadding), arg0) + return &MockSentPacketHandlerSetLastDatagramPaddingCall{Call: call} +} + +// MockSentPacketHandlerSetLastDatagramPaddingCall wrap *gomock.Call +type MockSentPacketHandlerSetLastDatagramPaddingCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerSetLastDatagramPaddingCall) Return() *MockSentPacketHandlerSetLastDatagramPaddingCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerSetLastDatagramPaddingCall) Do(f func(protocol.ByteCount)) *MockSentPacketHandlerSetLastDatagramPaddingCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerSetLastDatagramPaddingCall) DoAndReturn(f func(protocol.ByteCount)) *MockSentPacketHandlerSetLastDatagramPaddingCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetMaxDatagramSize mocks base method. +func (m *MockSentPacketHandler) SetMaxDatagramSize(count protocol.ByteCount) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetMaxDatagramSize", count) +} + +// SetMaxDatagramSize indicates an expected call of SetMaxDatagramSize. +func (mr *MockSentPacketHandlerMockRecorder) SetMaxDatagramSize(count any) *MockSentPacketHandlerSetMaxDatagramSizeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetMaxDatagramSize", reflect.TypeOf((*MockSentPacketHandler)(nil).SetMaxDatagramSize), count) + return &MockSentPacketHandlerSetMaxDatagramSizeCall{Call: call} +} + +// MockSentPacketHandlerSetMaxDatagramSizeCall wrap *gomock.Call +type MockSentPacketHandlerSetMaxDatagramSizeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerSetMaxDatagramSizeCall) Return() *MockSentPacketHandlerSetMaxDatagramSizeCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerSetMaxDatagramSizeCall) Do(f func(protocol.ByteCount)) *MockSentPacketHandlerSetMaxDatagramSizeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerSetMaxDatagramSizeCall) DoAndReturn(f func(protocol.ByteCount)) *MockSentPacketHandlerSetMaxDatagramSizeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// TimeUntilSend mocks base method. +func (m *MockSentPacketHandler) TimeUntilSend() monotime.Time { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "TimeUntilSend") + ret0, _ := ret[0].(monotime.Time) + return ret0 +} + +// TimeUntilSend indicates an expected call of TimeUntilSend. +func (mr *MockSentPacketHandlerMockRecorder) TimeUntilSend() *MockSentPacketHandlerTimeUntilSendCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "TimeUntilSend", reflect.TypeOf((*MockSentPacketHandler)(nil).TimeUntilSend)) + return &MockSentPacketHandlerTimeUntilSendCall{Call: call} +} + +// MockSentPacketHandlerTimeUntilSendCall wrap *gomock.Call +type MockSentPacketHandlerTimeUntilSendCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSentPacketHandlerTimeUntilSendCall) Return(arg0 monotime.Time) *MockSentPacketHandlerTimeUntilSendCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSentPacketHandlerTimeUntilSendCall) Do(f func() monotime.Time) *MockSentPacketHandlerTimeUntilSendCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSentPacketHandlerTimeUntilSendCall) DoAndReturn(f func() monotime.Time) *MockSentPacketHandlerTimeUntilSendCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/mocks/congestion.go b/third_party/quic-go/internal/mocks/congestion.go new file mode 100644 index 0000000..22cf0a2 --- /dev/null +++ b/third_party/quic-go/internal/mocks/congestion.go @@ -0,0 +1,486 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/congestion (interfaces: SendAlgorithmWithDebugInfos) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package mocks -destination congestion.go github.com/apernet/quic-go/internal/congestion SendAlgorithmWithDebugInfos +// + +// Package mocks is a generated GoMock package. +package mocks + +import ( + reflect "reflect" + + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockSendAlgorithmWithDebugInfos is a mock of SendAlgorithmWithDebugInfos interface. +type MockSendAlgorithmWithDebugInfos struct { + ctrl *gomock.Controller + recorder *MockSendAlgorithmWithDebugInfosMockRecorder + isgomock struct{} +} + +// MockSendAlgorithmWithDebugInfosMockRecorder is the mock recorder for MockSendAlgorithmWithDebugInfos. +type MockSendAlgorithmWithDebugInfosMockRecorder struct { + mock *MockSendAlgorithmWithDebugInfos +} + +// NewMockSendAlgorithmWithDebugInfos creates a new mock instance. +func NewMockSendAlgorithmWithDebugInfos(ctrl *gomock.Controller) *MockSendAlgorithmWithDebugInfos { + mock := &MockSendAlgorithmWithDebugInfos{ctrl: ctrl} + mock.recorder = &MockSendAlgorithmWithDebugInfosMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockSendAlgorithmWithDebugInfos) EXPECT() *MockSendAlgorithmWithDebugInfosMockRecorder { + return m.recorder +} + +// CanSend mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) CanSend(bytesInFlight protocol.ByteCount) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "CanSend", bytesInFlight) + ret0, _ := ret[0].(bool) + return ret0 +} + +// CanSend indicates an expected call of CanSend. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) CanSend(bytesInFlight any) *MockSendAlgorithmWithDebugInfosCanSendCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CanSend", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).CanSend), bytesInFlight) + return &MockSendAlgorithmWithDebugInfosCanSendCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosCanSendCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosCanSendCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosCanSendCall) Return(arg0 bool) *MockSendAlgorithmWithDebugInfosCanSendCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosCanSendCall) Do(f func(protocol.ByteCount) bool) *MockSendAlgorithmWithDebugInfosCanSendCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosCanSendCall) DoAndReturn(f func(protocol.ByteCount) bool) *MockSendAlgorithmWithDebugInfosCanSendCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetCongestionWindow mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) GetCongestionWindow() protocol.ByteCount { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetCongestionWindow") + ret0, _ := ret[0].(protocol.ByteCount) + return ret0 +} + +// GetCongestionWindow indicates an expected call of GetCongestionWindow. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) GetCongestionWindow() *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetCongestionWindow", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).GetCongestionWindow)) + return &MockSendAlgorithmWithDebugInfosGetCongestionWindowCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosGetCongestionWindowCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosGetCongestionWindowCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall) Return(arg0 protocol.ByteCount) *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall) Do(f func() protocol.ByteCount) *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall) DoAndReturn(f func() protocol.ByteCount) *MockSendAlgorithmWithDebugInfosGetCongestionWindowCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// HasPacingBudget mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) HasPacingBudget(now monotime.Time) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "HasPacingBudget", now) + ret0, _ := ret[0].(bool) + return ret0 +} + +// HasPacingBudget indicates an expected call of HasPacingBudget. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) HasPacingBudget(now any) *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasPacingBudget", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).HasPacingBudget), now) + return &MockSendAlgorithmWithDebugInfosHasPacingBudgetCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosHasPacingBudgetCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosHasPacingBudgetCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall) Return(arg0 bool) *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall) Do(f func(monotime.Time) bool) *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall) DoAndReturn(f func(monotime.Time) bool) *MockSendAlgorithmWithDebugInfosHasPacingBudgetCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// InRecovery mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) InRecovery() bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "InRecovery") + ret0, _ := ret[0].(bool) + return ret0 +} + +// InRecovery indicates an expected call of InRecovery. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) InRecovery() *MockSendAlgorithmWithDebugInfosInRecoveryCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "InRecovery", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).InRecovery)) + return &MockSendAlgorithmWithDebugInfosInRecoveryCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosInRecoveryCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosInRecoveryCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosInRecoveryCall) Return(arg0 bool) *MockSendAlgorithmWithDebugInfosInRecoveryCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosInRecoveryCall) Do(f func() bool) *MockSendAlgorithmWithDebugInfosInRecoveryCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosInRecoveryCall) DoAndReturn(f func() bool) *MockSendAlgorithmWithDebugInfosInRecoveryCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// InSlowStart mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) InSlowStart() bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "InSlowStart") + ret0, _ := ret[0].(bool) + return ret0 +} + +// InSlowStart indicates an expected call of InSlowStart. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) InSlowStart() *MockSendAlgorithmWithDebugInfosInSlowStartCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "InSlowStart", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).InSlowStart)) + return &MockSendAlgorithmWithDebugInfosInSlowStartCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosInSlowStartCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosInSlowStartCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosInSlowStartCall) Return(arg0 bool) *MockSendAlgorithmWithDebugInfosInSlowStartCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosInSlowStartCall) Do(f func() bool) *MockSendAlgorithmWithDebugInfosInSlowStartCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosInSlowStartCall) DoAndReturn(f func() bool) *MockSendAlgorithmWithDebugInfosInSlowStartCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// MaybeExitSlowStart mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) MaybeExitSlowStart() { + m.ctrl.T.Helper() + m.ctrl.Call(m, "MaybeExitSlowStart") +} + +// MaybeExitSlowStart indicates an expected call of MaybeExitSlowStart. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) MaybeExitSlowStart() *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "MaybeExitSlowStart", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).MaybeExitSlowStart)) + return &MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall) Return() *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall) Do(f func()) *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall) DoAndReturn(f func()) *MockSendAlgorithmWithDebugInfosMaybeExitSlowStartCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// OnCongestionEvent mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) OnCongestionEvent(number protocol.PacketNumber, lostBytes, priorInFlight protocol.ByteCount) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "OnCongestionEvent", number, lostBytes, priorInFlight) +} + +// OnCongestionEvent indicates an expected call of OnCongestionEvent. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) OnCongestionEvent(number, lostBytes, priorInFlight any) *MockSendAlgorithmWithDebugInfosOnCongestionEventCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnCongestionEvent", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).OnCongestionEvent), number, lostBytes, priorInFlight) + return &MockSendAlgorithmWithDebugInfosOnCongestionEventCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosOnCongestionEventCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosOnCongestionEventCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosOnCongestionEventCall) Return() *MockSendAlgorithmWithDebugInfosOnCongestionEventCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosOnCongestionEventCall) Do(f func(protocol.PacketNumber, protocol.ByteCount, protocol.ByteCount)) *MockSendAlgorithmWithDebugInfosOnCongestionEventCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosOnCongestionEventCall) DoAndReturn(f func(protocol.PacketNumber, protocol.ByteCount, protocol.ByteCount)) *MockSendAlgorithmWithDebugInfosOnCongestionEventCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// OnPacketAcked mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) OnPacketAcked(number protocol.PacketNumber, ackedBytes, priorInFlight protocol.ByteCount, eventTime monotime.Time) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "OnPacketAcked", number, ackedBytes, priorInFlight, eventTime) +} + +// OnPacketAcked indicates an expected call of OnPacketAcked. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) OnPacketAcked(number, ackedBytes, priorInFlight, eventTime any) *MockSendAlgorithmWithDebugInfosOnPacketAckedCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPacketAcked", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).OnPacketAcked), number, ackedBytes, priorInFlight, eventTime) + return &MockSendAlgorithmWithDebugInfosOnPacketAckedCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosOnPacketAckedCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosOnPacketAckedCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosOnPacketAckedCall) Return() *MockSendAlgorithmWithDebugInfosOnPacketAckedCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosOnPacketAckedCall) Do(f func(protocol.PacketNumber, protocol.ByteCount, protocol.ByteCount, monotime.Time)) *MockSendAlgorithmWithDebugInfosOnPacketAckedCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosOnPacketAckedCall) DoAndReturn(f func(protocol.PacketNumber, protocol.ByteCount, protocol.ByteCount, monotime.Time)) *MockSendAlgorithmWithDebugInfosOnPacketAckedCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// OnPacketSent mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) OnPacketSent(sentTime monotime.Time, bytesInFlight protocol.ByteCount, packetNumber protocol.PacketNumber, bytes protocol.ByteCount, isRetransmittable bool) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "OnPacketSent", sentTime, bytesInFlight, packetNumber, bytes, isRetransmittable) +} + +// OnPacketSent indicates an expected call of OnPacketSent. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) OnPacketSent(sentTime, bytesInFlight, packetNumber, bytes, isRetransmittable any) *MockSendAlgorithmWithDebugInfosOnPacketSentCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPacketSent", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).OnPacketSent), sentTime, bytesInFlight, packetNumber, bytes, isRetransmittable) + return &MockSendAlgorithmWithDebugInfosOnPacketSentCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosOnPacketSentCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosOnPacketSentCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosOnPacketSentCall) Return() *MockSendAlgorithmWithDebugInfosOnPacketSentCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosOnPacketSentCall) Do(f func(monotime.Time, protocol.ByteCount, protocol.PacketNumber, protocol.ByteCount, bool)) *MockSendAlgorithmWithDebugInfosOnPacketSentCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosOnPacketSentCall) DoAndReturn(f func(monotime.Time, protocol.ByteCount, protocol.PacketNumber, protocol.ByteCount, bool)) *MockSendAlgorithmWithDebugInfosOnPacketSentCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// OnRetransmissionTimeout mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) OnRetransmissionTimeout(packetsRetransmitted bool) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "OnRetransmissionTimeout", packetsRetransmitted) +} + +// OnRetransmissionTimeout indicates an expected call of OnRetransmissionTimeout. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) OnRetransmissionTimeout(packetsRetransmitted any) *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnRetransmissionTimeout", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).OnRetransmissionTimeout), packetsRetransmitted) + return &MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall) Return() *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall) Do(f func(bool)) *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall) DoAndReturn(f func(bool)) *MockSendAlgorithmWithDebugInfosOnRetransmissionTimeoutCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetMaxDatagramSize mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) SetMaxDatagramSize(arg0 protocol.ByteCount) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetMaxDatagramSize", arg0) +} + +// SetMaxDatagramSize indicates an expected call of SetMaxDatagramSize. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) SetMaxDatagramSize(arg0 any) *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetMaxDatagramSize", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).SetMaxDatagramSize), arg0) + return &MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall) Return() *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall) Do(f func(protocol.ByteCount)) *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall) DoAndReturn(f func(protocol.ByteCount)) *MockSendAlgorithmWithDebugInfosSetMaxDatagramSizeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// TimeUntilSend mocks base method. +func (m *MockSendAlgorithmWithDebugInfos) TimeUntilSend(bytesInFlight protocol.ByteCount) monotime.Time { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "TimeUntilSend", bytesInFlight) + ret0, _ := ret[0].(monotime.Time) + return ret0 +} + +// TimeUntilSend indicates an expected call of TimeUntilSend. +func (mr *MockSendAlgorithmWithDebugInfosMockRecorder) TimeUntilSend(bytesInFlight any) *MockSendAlgorithmWithDebugInfosTimeUntilSendCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "TimeUntilSend", reflect.TypeOf((*MockSendAlgorithmWithDebugInfos)(nil).TimeUntilSend), bytesInFlight) + return &MockSendAlgorithmWithDebugInfosTimeUntilSendCall{Call: call} +} + +// MockSendAlgorithmWithDebugInfosTimeUntilSendCall wrap *gomock.Call +type MockSendAlgorithmWithDebugInfosTimeUntilSendCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendAlgorithmWithDebugInfosTimeUntilSendCall) Return(arg0 monotime.Time) *MockSendAlgorithmWithDebugInfosTimeUntilSendCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendAlgorithmWithDebugInfosTimeUntilSendCall) Do(f func(protocol.ByteCount) monotime.Time) *MockSendAlgorithmWithDebugInfosTimeUntilSendCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendAlgorithmWithDebugInfosTimeUntilSendCall) DoAndReturn(f func(protocol.ByteCount) monotime.Time) *MockSendAlgorithmWithDebugInfosTimeUntilSendCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/mocks/crypto_setup.go b/third_party/quic-go/internal/mocks/crypto_setup.go new file mode 100644 index 0000000..70d7305 --- /dev/null +++ b/third_party/quic-go/internal/mocks/crypto_setup.go @@ -0,0 +1,730 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/handshake (interfaces: CryptoSetup) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package mocks -destination crypto_setup.go github.com/apernet/quic-go/internal/handshake CryptoSetup +// + +// Package mocks is a generated GoMock package. +package mocks + +import ( + context "context" + reflect "reflect" + + handshake "github.com/apernet/quic-go/internal/handshake" + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockCryptoSetup is a mock of CryptoSetup interface. +type MockCryptoSetup struct { + ctrl *gomock.Controller + recorder *MockCryptoSetupMockRecorder + isgomock struct{} +} + +// MockCryptoSetupMockRecorder is the mock recorder for MockCryptoSetup. +type MockCryptoSetupMockRecorder struct { + mock *MockCryptoSetup +} + +// NewMockCryptoSetup creates a new mock instance. +func NewMockCryptoSetup(ctrl *gomock.Controller) *MockCryptoSetup { + mock := &MockCryptoSetup{ctrl: ctrl} + mock.recorder = &MockCryptoSetupMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockCryptoSetup) EXPECT() *MockCryptoSetupMockRecorder { + return m.recorder +} + +// ChangeConnectionID mocks base method. +func (m *MockCryptoSetup) ChangeConnectionID(arg0 protocol.ConnectionID) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "ChangeConnectionID", arg0) +} + +// ChangeConnectionID indicates an expected call of ChangeConnectionID. +func (mr *MockCryptoSetupMockRecorder) ChangeConnectionID(arg0 any) *MockCryptoSetupChangeConnectionIDCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ChangeConnectionID", reflect.TypeOf((*MockCryptoSetup)(nil).ChangeConnectionID), arg0) + return &MockCryptoSetupChangeConnectionIDCall{Call: call} +} + +// MockCryptoSetupChangeConnectionIDCall wrap *gomock.Call +type MockCryptoSetupChangeConnectionIDCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupChangeConnectionIDCall) Return() *MockCryptoSetupChangeConnectionIDCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupChangeConnectionIDCall) Do(f func(protocol.ConnectionID)) *MockCryptoSetupChangeConnectionIDCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupChangeConnectionIDCall) DoAndReturn(f func(protocol.ConnectionID)) *MockCryptoSetupChangeConnectionIDCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Close mocks base method. +func (m *MockCryptoSetup) Close() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Close") + ret0, _ := ret[0].(error) + return ret0 +} + +// Close indicates an expected call of Close. +func (mr *MockCryptoSetupMockRecorder) Close() *MockCryptoSetupCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockCryptoSetup)(nil).Close)) + return &MockCryptoSetupCloseCall{Call: call} +} + +// MockCryptoSetupCloseCall wrap *gomock.Call +type MockCryptoSetupCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupCloseCall) Return(arg0 error) *MockCryptoSetupCloseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupCloseCall) Do(f func() error) *MockCryptoSetupCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupCloseCall) DoAndReturn(f func() error) *MockCryptoSetupCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ConnectionState mocks base method. +func (m *MockCryptoSetup) ConnectionState() handshake.ConnectionState { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ConnectionState") + ret0, _ := ret[0].(handshake.ConnectionState) + return ret0 +} + +// ConnectionState indicates an expected call of ConnectionState. +func (mr *MockCryptoSetupMockRecorder) ConnectionState() *MockCryptoSetupConnectionStateCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ConnectionState", reflect.TypeOf((*MockCryptoSetup)(nil).ConnectionState)) + return &MockCryptoSetupConnectionStateCall{Call: call} +} + +// MockCryptoSetupConnectionStateCall wrap *gomock.Call +type MockCryptoSetupConnectionStateCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupConnectionStateCall) Return(arg0 handshake.ConnectionState) *MockCryptoSetupConnectionStateCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupConnectionStateCall) Do(f func() handshake.ConnectionState) *MockCryptoSetupConnectionStateCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupConnectionStateCall) DoAndReturn(f func() handshake.ConnectionState) *MockCryptoSetupConnectionStateCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// DiscardInitialKeys mocks base method. +func (m *MockCryptoSetup) DiscardInitialKeys() { + m.ctrl.T.Helper() + m.ctrl.Call(m, "DiscardInitialKeys") +} + +// DiscardInitialKeys indicates an expected call of DiscardInitialKeys. +func (mr *MockCryptoSetupMockRecorder) DiscardInitialKeys() *MockCryptoSetupDiscardInitialKeysCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DiscardInitialKeys", reflect.TypeOf((*MockCryptoSetup)(nil).DiscardInitialKeys)) + return &MockCryptoSetupDiscardInitialKeysCall{Call: call} +} + +// MockCryptoSetupDiscardInitialKeysCall wrap *gomock.Call +type MockCryptoSetupDiscardInitialKeysCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupDiscardInitialKeysCall) Return() *MockCryptoSetupDiscardInitialKeysCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupDiscardInitialKeysCall) Do(f func()) *MockCryptoSetupDiscardInitialKeysCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupDiscardInitialKeysCall) DoAndReturn(f func()) *MockCryptoSetupDiscardInitialKeysCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Get0RTTOpener mocks base method. +func (m *MockCryptoSetup) Get0RTTOpener() (handshake.LongHeaderOpener, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Get0RTTOpener") + ret0, _ := ret[0].(handshake.LongHeaderOpener) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Get0RTTOpener indicates an expected call of Get0RTTOpener. +func (mr *MockCryptoSetupMockRecorder) Get0RTTOpener() *MockCryptoSetupGet0RTTOpenerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get0RTTOpener", reflect.TypeOf((*MockCryptoSetup)(nil).Get0RTTOpener)) + return &MockCryptoSetupGet0RTTOpenerCall{Call: call} +} + +// MockCryptoSetupGet0RTTOpenerCall wrap *gomock.Call +type MockCryptoSetupGet0RTTOpenerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGet0RTTOpenerCall) Return(arg0 handshake.LongHeaderOpener, arg1 error) *MockCryptoSetupGet0RTTOpenerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGet0RTTOpenerCall) Do(f func() (handshake.LongHeaderOpener, error)) *MockCryptoSetupGet0RTTOpenerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGet0RTTOpenerCall) DoAndReturn(f func() (handshake.LongHeaderOpener, error)) *MockCryptoSetupGet0RTTOpenerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Get0RTTSealer mocks base method. +func (m *MockCryptoSetup) Get0RTTSealer() (handshake.LongHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Get0RTTSealer") + ret0, _ := ret[0].(handshake.LongHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Get0RTTSealer indicates an expected call of Get0RTTSealer. +func (mr *MockCryptoSetupMockRecorder) Get0RTTSealer() *MockCryptoSetupGet0RTTSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get0RTTSealer", reflect.TypeOf((*MockCryptoSetup)(nil).Get0RTTSealer)) + return &MockCryptoSetupGet0RTTSealerCall{Call: call} +} + +// MockCryptoSetupGet0RTTSealerCall wrap *gomock.Call +type MockCryptoSetupGet0RTTSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGet0RTTSealerCall) Return(arg0 handshake.LongHeaderSealer, arg1 error) *MockCryptoSetupGet0RTTSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGet0RTTSealerCall) Do(f func() (handshake.LongHeaderSealer, error)) *MockCryptoSetupGet0RTTSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGet0RTTSealerCall) DoAndReturn(f func() (handshake.LongHeaderSealer, error)) *MockCryptoSetupGet0RTTSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Get1RTTOpener mocks base method. +func (m *MockCryptoSetup) Get1RTTOpener() (handshake.ShortHeaderOpener, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Get1RTTOpener") + ret0, _ := ret[0].(handshake.ShortHeaderOpener) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Get1RTTOpener indicates an expected call of Get1RTTOpener. +func (mr *MockCryptoSetupMockRecorder) Get1RTTOpener() *MockCryptoSetupGet1RTTOpenerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get1RTTOpener", reflect.TypeOf((*MockCryptoSetup)(nil).Get1RTTOpener)) + return &MockCryptoSetupGet1RTTOpenerCall{Call: call} +} + +// MockCryptoSetupGet1RTTOpenerCall wrap *gomock.Call +type MockCryptoSetupGet1RTTOpenerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGet1RTTOpenerCall) Return(arg0 handshake.ShortHeaderOpener, arg1 error) *MockCryptoSetupGet1RTTOpenerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGet1RTTOpenerCall) Do(f func() (handshake.ShortHeaderOpener, error)) *MockCryptoSetupGet1RTTOpenerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGet1RTTOpenerCall) DoAndReturn(f func() (handshake.ShortHeaderOpener, error)) *MockCryptoSetupGet1RTTOpenerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Get1RTTSealer mocks base method. +func (m *MockCryptoSetup) Get1RTTSealer() (handshake.ShortHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Get1RTTSealer") + ret0, _ := ret[0].(handshake.ShortHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Get1RTTSealer indicates an expected call of Get1RTTSealer. +func (mr *MockCryptoSetupMockRecorder) Get1RTTSealer() *MockCryptoSetupGet1RTTSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get1RTTSealer", reflect.TypeOf((*MockCryptoSetup)(nil).Get1RTTSealer)) + return &MockCryptoSetupGet1RTTSealerCall{Call: call} +} + +// MockCryptoSetupGet1RTTSealerCall wrap *gomock.Call +type MockCryptoSetupGet1RTTSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGet1RTTSealerCall) Return(arg0 handshake.ShortHeaderSealer, arg1 error) *MockCryptoSetupGet1RTTSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGet1RTTSealerCall) Do(f func() (handshake.ShortHeaderSealer, error)) *MockCryptoSetupGet1RTTSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGet1RTTSealerCall) DoAndReturn(f func() (handshake.ShortHeaderSealer, error)) *MockCryptoSetupGet1RTTSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetHandshakeOpener mocks base method. +func (m *MockCryptoSetup) GetHandshakeOpener() (handshake.LongHeaderOpener, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetHandshakeOpener") + ret0, _ := ret[0].(handshake.LongHeaderOpener) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetHandshakeOpener indicates an expected call of GetHandshakeOpener. +func (mr *MockCryptoSetupMockRecorder) GetHandshakeOpener() *MockCryptoSetupGetHandshakeOpenerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetHandshakeOpener", reflect.TypeOf((*MockCryptoSetup)(nil).GetHandshakeOpener)) + return &MockCryptoSetupGetHandshakeOpenerCall{Call: call} +} + +// MockCryptoSetupGetHandshakeOpenerCall wrap *gomock.Call +type MockCryptoSetupGetHandshakeOpenerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGetHandshakeOpenerCall) Return(arg0 handshake.LongHeaderOpener, arg1 error) *MockCryptoSetupGetHandshakeOpenerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGetHandshakeOpenerCall) Do(f func() (handshake.LongHeaderOpener, error)) *MockCryptoSetupGetHandshakeOpenerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGetHandshakeOpenerCall) DoAndReturn(f func() (handshake.LongHeaderOpener, error)) *MockCryptoSetupGetHandshakeOpenerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetHandshakeSealer mocks base method. +func (m *MockCryptoSetup) GetHandshakeSealer() (handshake.LongHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetHandshakeSealer") + ret0, _ := ret[0].(handshake.LongHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetHandshakeSealer indicates an expected call of GetHandshakeSealer. +func (mr *MockCryptoSetupMockRecorder) GetHandshakeSealer() *MockCryptoSetupGetHandshakeSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetHandshakeSealer", reflect.TypeOf((*MockCryptoSetup)(nil).GetHandshakeSealer)) + return &MockCryptoSetupGetHandshakeSealerCall{Call: call} +} + +// MockCryptoSetupGetHandshakeSealerCall wrap *gomock.Call +type MockCryptoSetupGetHandshakeSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGetHandshakeSealerCall) Return(arg0 handshake.LongHeaderSealer, arg1 error) *MockCryptoSetupGetHandshakeSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGetHandshakeSealerCall) Do(f func() (handshake.LongHeaderSealer, error)) *MockCryptoSetupGetHandshakeSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGetHandshakeSealerCall) DoAndReturn(f func() (handshake.LongHeaderSealer, error)) *MockCryptoSetupGetHandshakeSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetInitialOpener mocks base method. +func (m *MockCryptoSetup) GetInitialOpener() (handshake.LongHeaderOpener, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetInitialOpener") + ret0, _ := ret[0].(handshake.LongHeaderOpener) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetInitialOpener indicates an expected call of GetInitialOpener. +func (mr *MockCryptoSetupMockRecorder) GetInitialOpener() *MockCryptoSetupGetInitialOpenerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetInitialOpener", reflect.TypeOf((*MockCryptoSetup)(nil).GetInitialOpener)) + return &MockCryptoSetupGetInitialOpenerCall{Call: call} +} + +// MockCryptoSetupGetInitialOpenerCall wrap *gomock.Call +type MockCryptoSetupGetInitialOpenerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGetInitialOpenerCall) Return(arg0 handshake.LongHeaderOpener, arg1 error) *MockCryptoSetupGetInitialOpenerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGetInitialOpenerCall) Do(f func() (handshake.LongHeaderOpener, error)) *MockCryptoSetupGetInitialOpenerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGetInitialOpenerCall) DoAndReturn(f func() (handshake.LongHeaderOpener, error)) *MockCryptoSetupGetInitialOpenerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetInitialSealer mocks base method. +func (m *MockCryptoSetup) GetInitialSealer() (handshake.LongHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetInitialSealer") + ret0, _ := ret[0].(handshake.LongHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetInitialSealer indicates an expected call of GetInitialSealer. +func (mr *MockCryptoSetupMockRecorder) GetInitialSealer() *MockCryptoSetupGetInitialSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetInitialSealer", reflect.TypeOf((*MockCryptoSetup)(nil).GetInitialSealer)) + return &MockCryptoSetupGetInitialSealerCall{Call: call} +} + +// MockCryptoSetupGetInitialSealerCall wrap *gomock.Call +type MockCryptoSetupGetInitialSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGetInitialSealerCall) Return(arg0 handshake.LongHeaderSealer, arg1 error) *MockCryptoSetupGetInitialSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGetInitialSealerCall) Do(f func() (handshake.LongHeaderSealer, error)) *MockCryptoSetupGetInitialSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGetInitialSealerCall) DoAndReturn(f func() (handshake.LongHeaderSealer, error)) *MockCryptoSetupGetInitialSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetSessionTicket mocks base method. +func (m *MockCryptoSetup) GetSessionTicket() ([]byte, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetSessionTicket") + ret0, _ := ret[0].([]byte) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetSessionTicket indicates an expected call of GetSessionTicket. +func (mr *MockCryptoSetupMockRecorder) GetSessionTicket() *MockCryptoSetupGetSessionTicketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetSessionTicket", reflect.TypeOf((*MockCryptoSetup)(nil).GetSessionTicket)) + return &MockCryptoSetupGetSessionTicketCall{Call: call} +} + +// MockCryptoSetupGetSessionTicketCall wrap *gomock.Call +type MockCryptoSetupGetSessionTicketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupGetSessionTicketCall) Return(arg0 []byte, arg1 error) *MockCryptoSetupGetSessionTicketCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupGetSessionTicketCall) Do(f func() ([]byte, error)) *MockCryptoSetupGetSessionTicketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupGetSessionTicketCall) DoAndReturn(f func() ([]byte, error)) *MockCryptoSetupGetSessionTicketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// HandleMessage mocks base method. +func (m *MockCryptoSetup) HandleMessage(arg0 []byte, arg1 protocol.EncryptionLevel) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "HandleMessage", arg0, arg1) + ret0, _ := ret[0].(error) + return ret0 +} + +// HandleMessage indicates an expected call of HandleMessage. +func (mr *MockCryptoSetupMockRecorder) HandleMessage(arg0, arg1 any) *MockCryptoSetupHandleMessageCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HandleMessage", reflect.TypeOf((*MockCryptoSetup)(nil).HandleMessage), arg0, arg1) + return &MockCryptoSetupHandleMessageCall{Call: call} +} + +// MockCryptoSetupHandleMessageCall wrap *gomock.Call +type MockCryptoSetupHandleMessageCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupHandleMessageCall) Return(arg0 error) *MockCryptoSetupHandleMessageCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupHandleMessageCall) Do(f func([]byte, protocol.EncryptionLevel) error) *MockCryptoSetupHandleMessageCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupHandleMessageCall) DoAndReturn(f func([]byte, protocol.EncryptionLevel) error) *MockCryptoSetupHandleMessageCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// NextEvent mocks base method. +func (m *MockCryptoSetup) NextEvent() handshake.Event { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "NextEvent") + ret0, _ := ret[0].(handshake.Event) + return ret0 +} + +// NextEvent indicates an expected call of NextEvent. +func (mr *MockCryptoSetupMockRecorder) NextEvent() *MockCryptoSetupNextEventCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "NextEvent", reflect.TypeOf((*MockCryptoSetup)(nil).NextEvent)) + return &MockCryptoSetupNextEventCall{Call: call} +} + +// MockCryptoSetupNextEventCall wrap *gomock.Call +type MockCryptoSetupNextEventCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupNextEventCall) Return(arg0 handshake.Event) *MockCryptoSetupNextEventCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupNextEventCall) Do(f func() handshake.Event) *MockCryptoSetupNextEventCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupNextEventCall) DoAndReturn(f func() handshake.Event) *MockCryptoSetupNextEventCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetHandshakeConfirmed mocks base method. +func (m *MockCryptoSetup) SetHandshakeConfirmed() { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetHandshakeConfirmed") +} + +// SetHandshakeConfirmed indicates an expected call of SetHandshakeConfirmed. +func (mr *MockCryptoSetupMockRecorder) SetHandshakeConfirmed() *MockCryptoSetupSetHandshakeConfirmedCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetHandshakeConfirmed", reflect.TypeOf((*MockCryptoSetup)(nil).SetHandshakeConfirmed)) + return &MockCryptoSetupSetHandshakeConfirmedCall{Call: call} +} + +// MockCryptoSetupSetHandshakeConfirmedCall wrap *gomock.Call +type MockCryptoSetupSetHandshakeConfirmedCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupSetHandshakeConfirmedCall) Return() *MockCryptoSetupSetHandshakeConfirmedCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupSetHandshakeConfirmedCall) Do(f func()) *MockCryptoSetupSetHandshakeConfirmedCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupSetHandshakeConfirmedCall) DoAndReturn(f func()) *MockCryptoSetupSetHandshakeConfirmedCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetLargest1RTTAcked mocks base method. +func (m *MockCryptoSetup) SetLargest1RTTAcked(arg0 protocol.PacketNumber) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetLargest1RTTAcked", arg0) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetLargest1RTTAcked indicates an expected call of SetLargest1RTTAcked. +func (mr *MockCryptoSetupMockRecorder) SetLargest1RTTAcked(arg0 any) *MockCryptoSetupSetLargest1RTTAckedCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetLargest1RTTAcked", reflect.TypeOf((*MockCryptoSetup)(nil).SetLargest1RTTAcked), arg0) + return &MockCryptoSetupSetLargest1RTTAckedCall{Call: call} +} + +// MockCryptoSetupSetLargest1RTTAckedCall wrap *gomock.Call +type MockCryptoSetupSetLargest1RTTAckedCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupSetLargest1RTTAckedCall) Return(arg0 error) *MockCryptoSetupSetLargest1RTTAckedCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupSetLargest1RTTAckedCall) Do(f func(protocol.PacketNumber) error) *MockCryptoSetupSetLargest1RTTAckedCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupSetLargest1RTTAckedCall) DoAndReturn(f func(protocol.PacketNumber) error) *MockCryptoSetupSetLargest1RTTAckedCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// StartHandshake mocks base method. +func (m *MockCryptoSetup) StartHandshake(arg0 context.Context) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "StartHandshake", arg0) + ret0, _ := ret[0].(error) + return ret0 +} + +// StartHandshake indicates an expected call of StartHandshake. +func (mr *MockCryptoSetupMockRecorder) StartHandshake(arg0 any) *MockCryptoSetupStartHandshakeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StartHandshake", reflect.TypeOf((*MockCryptoSetup)(nil).StartHandshake), arg0) + return &MockCryptoSetupStartHandshakeCall{Call: call} +} + +// MockCryptoSetupStartHandshakeCall wrap *gomock.Call +type MockCryptoSetupStartHandshakeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockCryptoSetupStartHandshakeCall) Return(arg0 error) *MockCryptoSetupStartHandshakeCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockCryptoSetupStartHandshakeCall) Do(f func(context.Context) error) *MockCryptoSetupStartHandshakeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockCryptoSetupStartHandshakeCall) DoAndReturn(f func(context.Context) error) *MockCryptoSetupStartHandshakeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/mocks/long_header_opener.go b/third_party/quic-go/internal/mocks/long_header_opener.go new file mode 100644 index 0000000..b3f220d --- /dev/null +++ b/third_party/quic-go/internal/mocks/long_header_opener.go @@ -0,0 +1,154 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/handshake (interfaces: LongHeaderOpener) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package mocks -destination long_header_opener.go github.com/apernet/quic-go/internal/handshake LongHeaderOpener +// + +// Package mocks is a generated GoMock package. +package mocks + +import ( + reflect "reflect" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockLongHeaderOpener is a mock of LongHeaderOpener interface. +type MockLongHeaderOpener struct { + ctrl *gomock.Controller + recorder *MockLongHeaderOpenerMockRecorder + isgomock struct{} +} + +// MockLongHeaderOpenerMockRecorder is the mock recorder for MockLongHeaderOpener. +type MockLongHeaderOpenerMockRecorder struct { + mock *MockLongHeaderOpener +} + +// NewMockLongHeaderOpener creates a new mock instance. +func NewMockLongHeaderOpener(ctrl *gomock.Controller) *MockLongHeaderOpener { + mock := &MockLongHeaderOpener{ctrl: ctrl} + mock.recorder = &MockLongHeaderOpenerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockLongHeaderOpener) EXPECT() *MockLongHeaderOpenerMockRecorder { + return m.recorder +} + +// DecodePacketNumber mocks base method. +func (m *MockLongHeaderOpener) DecodePacketNumber(wirePN protocol.PacketNumber, wirePNLen protocol.PacketNumberLen) protocol.PacketNumber { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DecodePacketNumber", wirePN, wirePNLen) + ret0, _ := ret[0].(protocol.PacketNumber) + return ret0 +} + +// DecodePacketNumber indicates an expected call of DecodePacketNumber. +func (mr *MockLongHeaderOpenerMockRecorder) DecodePacketNumber(wirePN, wirePNLen any) *MockLongHeaderOpenerDecodePacketNumberCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DecodePacketNumber", reflect.TypeOf((*MockLongHeaderOpener)(nil).DecodePacketNumber), wirePN, wirePNLen) + return &MockLongHeaderOpenerDecodePacketNumberCall{Call: call} +} + +// MockLongHeaderOpenerDecodePacketNumberCall wrap *gomock.Call +type MockLongHeaderOpenerDecodePacketNumberCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockLongHeaderOpenerDecodePacketNumberCall) Return(arg0 protocol.PacketNumber) *MockLongHeaderOpenerDecodePacketNumberCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockLongHeaderOpenerDecodePacketNumberCall) Do(f func(protocol.PacketNumber, protocol.PacketNumberLen) protocol.PacketNumber) *MockLongHeaderOpenerDecodePacketNumberCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockLongHeaderOpenerDecodePacketNumberCall) DoAndReturn(f func(protocol.PacketNumber, protocol.PacketNumberLen) protocol.PacketNumber) *MockLongHeaderOpenerDecodePacketNumberCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// DecryptHeader mocks base method. +func (m *MockLongHeaderOpener) DecryptHeader(sample []byte, firstByte *byte, pnBytes []byte) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "DecryptHeader", sample, firstByte, pnBytes) +} + +// DecryptHeader indicates an expected call of DecryptHeader. +func (mr *MockLongHeaderOpenerMockRecorder) DecryptHeader(sample, firstByte, pnBytes any) *MockLongHeaderOpenerDecryptHeaderCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DecryptHeader", reflect.TypeOf((*MockLongHeaderOpener)(nil).DecryptHeader), sample, firstByte, pnBytes) + return &MockLongHeaderOpenerDecryptHeaderCall{Call: call} +} + +// MockLongHeaderOpenerDecryptHeaderCall wrap *gomock.Call +type MockLongHeaderOpenerDecryptHeaderCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockLongHeaderOpenerDecryptHeaderCall) Return() *MockLongHeaderOpenerDecryptHeaderCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockLongHeaderOpenerDecryptHeaderCall) Do(f func([]byte, *byte, []byte)) *MockLongHeaderOpenerDecryptHeaderCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockLongHeaderOpenerDecryptHeaderCall) DoAndReturn(f func([]byte, *byte, []byte)) *MockLongHeaderOpenerDecryptHeaderCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Open mocks base method. +func (m *MockLongHeaderOpener) Open(dst, src []byte, pn protocol.PacketNumber, associatedData []byte) ([]byte, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Open", dst, src, pn, associatedData) + ret0, _ := ret[0].([]byte) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Open indicates an expected call of Open. +func (mr *MockLongHeaderOpenerMockRecorder) Open(dst, src, pn, associatedData any) *MockLongHeaderOpenerOpenCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Open", reflect.TypeOf((*MockLongHeaderOpener)(nil).Open), dst, src, pn, associatedData) + return &MockLongHeaderOpenerOpenCall{Call: call} +} + +// MockLongHeaderOpenerOpenCall wrap *gomock.Call +type MockLongHeaderOpenerOpenCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockLongHeaderOpenerOpenCall) Return(arg0 []byte, arg1 error) *MockLongHeaderOpenerOpenCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockLongHeaderOpenerOpenCall) Do(f func([]byte, []byte, protocol.PacketNumber, []byte) ([]byte, error)) *MockLongHeaderOpenerOpenCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockLongHeaderOpenerOpenCall) DoAndReturn(f func([]byte, []byte, protocol.PacketNumber, []byte) ([]byte, error)) *MockLongHeaderOpenerOpenCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/mocks/mockgen.go b/third_party/quic-go/internal/mocks/mockgen.go new file mode 100644 index 0000000..1ac86ed --- /dev/null +++ b/third_party/quic-go/internal/mocks/mockgen.go @@ -0,0 +1,10 @@ +//go:build gomock || generate + +package mocks + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package mocks -destination short_header_sealer.go github.com/apernet/quic-go/internal/handshake ShortHeaderSealer" +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package mocks -destination short_header_opener.go github.com/apernet/quic-go/internal/handshake ShortHeaderOpener" +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package mocks -destination long_header_opener.go github.com/apernet/quic-go/internal/handshake LongHeaderOpener" +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package mocks -destination crypto_setup.go github.com/apernet/quic-go/internal/handshake CryptoSetup" +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package mocks -destination congestion.go github.com/apernet/quic-go/internal/congestion SendAlgorithmWithDebugInfos" +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package mockackhandler -destination ackhandler/sent_packet_handler.go github.com/apernet/quic-go/internal/ackhandler SentPacketHandler" diff --git a/third_party/quic-go/internal/mocks/short_header_opener.go b/third_party/quic-go/internal/mocks/short_header_opener.go new file mode 100644 index 0000000..56b9b2d --- /dev/null +++ b/third_party/quic-go/internal/mocks/short_header_opener.go @@ -0,0 +1,155 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/handshake (interfaces: ShortHeaderOpener) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package mocks -destination short_header_opener.go github.com/apernet/quic-go/internal/handshake ShortHeaderOpener +// + +// Package mocks is a generated GoMock package. +package mocks + +import ( + reflect "reflect" + + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockShortHeaderOpener is a mock of ShortHeaderOpener interface. +type MockShortHeaderOpener struct { + ctrl *gomock.Controller + recorder *MockShortHeaderOpenerMockRecorder + isgomock struct{} +} + +// MockShortHeaderOpenerMockRecorder is the mock recorder for MockShortHeaderOpener. +type MockShortHeaderOpenerMockRecorder struct { + mock *MockShortHeaderOpener +} + +// NewMockShortHeaderOpener creates a new mock instance. +func NewMockShortHeaderOpener(ctrl *gomock.Controller) *MockShortHeaderOpener { + mock := &MockShortHeaderOpener{ctrl: ctrl} + mock.recorder = &MockShortHeaderOpenerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockShortHeaderOpener) EXPECT() *MockShortHeaderOpenerMockRecorder { + return m.recorder +} + +// DecodePacketNumber mocks base method. +func (m *MockShortHeaderOpener) DecodePacketNumber(wirePN protocol.PacketNumber, wirePNLen protocol.PacketNumberLen) protocol.PacketNumber { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DecodePacketNumber", wirePN, wirePNLen) + ret0, _ := ret[0].(protocol.PacketNumber) + return ret0 +} + +// DecodePacketNumber indicates an expected call of DecodePacketNumber. +func (mr *MockShortHeaderOpenerMockRecorder) DecodePacketNumber(wirePN, wirePNLen any) *MockShortHeaderOpenerDecodePacketNumberCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DecodePacketNumber", reflect.TypeOf((*MockShortHeaderOpener)(nil).DecodePacketNumber), wirePN, wirePNLen) + return &MockShortHeaderOpenerDecodePacketNumberCall{Call: call} +} + +// MockShortHeaderOpenerDecodePacketNumberCall wrap *gomock.Call +type MockShortHeaderOpenerDecodePacketNumberCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderOpenerDecodePacketNumberCall) Return(arg0 protocol.PacketNumber) *MockShortHeaderOpenerDecodePacketNumberCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderOpenerDecodePacketNumberCall) Do(f func(protocol.PacketNumber, protocol.PacketNumberLen) protocol.PacketNumber) *MockShortHeaderOpenerDecodePacketNumberCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderOpenerDecodePacketNumberCall) DoAndReturn(f func(protocol.PacketNumber, protocol.PacketNumberLen) protocol.PacketNumber) *MockShortHeaderOpenerDecodePacketNumberCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// DecryptHeader mocks base method. +func (m *MockShortHeaderOpener) DecryptHeader(sample []byte, firstByte *byte, pnBytes []byte) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "DecryptHeader", sample, firstByte, pnBytes) +} + +// DecryptHeader indicates an expected call of DecryptHeader. +func (mr *MockShortHeaderOpenerMockRecorder) DecryptHeader(sample, firstByte, pnBytes any) *MockShortHeaderOpenerDecryptHeaderCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DecryptHeader", reflect.TypeOf((*MockShortHeaderOpener)(nil).DecryptHeader), sample, firstByte, pnBytes) + return &MockShortHeaderOpenerDecryptHeaderCall{Call: call} +} + +// MockShortHeaderOpenerDecryptHeaderCall wrap *gomock.Call +type MockShortHeaderOpenerDecryptHeaderCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderOpenerDecryptHeaderCall) Return() *MockShortHeaderOpenerDecryptHeaderCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderOpenerDecryptHeaderCall) Do(f func([]byte, *byte, []byte)) *MockShortHeaderOpenerDecryptHeaderCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderOpenerDecryptHeaderCall) DoAndReturn(f func([]byte, *byte, []byte)) *MockShortHeaderOpenerDecryptHeaderCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Open mocks base method. +func (m *MockShortHeaderOpener) Open(dst, src []byte, rcvTime monotime.Time, pn protocol.PacketNumber, kp protocol.KeyPhaseBit, associatedData []byte) ([]byte, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Open", dst, src, rcvTime, pn, kp, associatedData) + ret0, _ := ret[0].([]byte) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Open indicates an expected call of Open. +func (mr *MockShortHeaderOpenerMockRecorder) Open(dst, src, rcvTime, pn, kp, associatedData any) *MockShortHeaderOpenerOpenCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Open", reflect.TypeOf((*MockShortHeaderOpener)(nil).Open), dst, src, rcvTime, pn, kp, associatedData) + return &MockShortHeaderOpenerOpenCall{Call: call} +} + +// MockShortHeaderOpenerOpenCall wrap *gomock.Call +type MockShortHeaderOpenerOpenCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderOpenerOpenCall) Return(arg0 []byte, arg1 error) *MockShortHeaderOpenerOpenCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderOpenerOpenCall) Do(f func([]byte, []byte, monotime.Time, protocol.PacketNumber, protocol.KeyPhaseBit, []byte) ([]byte, error)) *MockShortHeaderOpenerOpenCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderOpenerOpenCall) DoAndReturn(f func([]byte, []byte, monotime.Time, protocol.PacketNumber, protocol.KeyPhaseBit, []byte) ([]byte, error)) *MockShortHeaderOpenerOpenCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/mocks/short_header_sealer.go b/third_party/quic-go/internal/mocks/short_header_sealer.go new file mode 100644 index 0000000..2f5562a --- /dev/null +++ b/third_party/quic-go/internal/mocks/short_header_sealer.go @@ -0,0 +1,191 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go/internal/handshake (interfaces: ShortHeaderSealer) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package mocks -destination short_header_sealer.go github.com/apernet/quic-go/internal/handshake ShortHeaderSealer +// + +// Package mocks is a generated GoMock package. +package mocks + +import ( + reflect "reflect" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockShortHeaderSealer is a mock of ShortHeaderSealer interface. +type MockShortHeaderSealer struct { + ctrl *gomock.Controller + recorder *MockShortHeaderSealerMockRecorder + isgomock struct{} +} + +// MockShortHeaderSealerMockRecorder is the mock recorder for MockShortHeaderSealer. +type MockShortHeaderSealerMockRecorder struct { + mock *MockShortHeaderSealer +} + +// NewMockShortHeaderSealer creates a new mock instance. +func NewMockShortHeaderSealer(ctrl *gomock.Controller) *MockShortHeaderSealer { + mock := &MockShortHeaderSealer{ctrl: ctrl} + mock.recorder = &MockShortHeaderSealerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockShortHeaderSealer) EXPECT() *MockShortHeaderSealerMockRecorder { + return m.recorder +} + +// EncryptHeader mocks base method. +func (m *MockShortHeaderSealer) EncryptHeader(sample []byte, firstByte *byte, pnBytes []byte) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "EncryptHeader", sample, firstByte, pnBytes) +} + +// EncryptHeader indicates an expected call of EncryptHeader. +func (mr *MockShortHeaderSealerMockRecorder) EncryptHeader(sample, firstByte, pnBytes any) *MockShortHeaderSealerEncryptHeaderCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "EncryptHeader", reflect.TypeOf((*MockShortHeaderSealer)(nil).EncryptHeader), sample, firstByte, pnBytes) + return &MockShortHeaderSealerEncryptHeaderCall{Call: call} +} + +// MockShortHeaderSealerEncryptHeaderCall wrap *gomock.Call +type MockShortHeaderSealerEncryptHeaderCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderSealerEncryptHeaderCall) Return() *MockShortHeaderSealerEncryptHeaderCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderSealerEncryptHeaderCall) Do(f func([]byte, *byte, []byte)) *MockShortHeaderSealerEncryptHeaderCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderSealerEncryptHeaderCall) DoAndReturn(f func([]byte, *byte, []byte)) *MockShortHeaderSealerEncryptHeaderCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// KeyPhase mocks base method. +func (m *MockShortHeaderSealer) KeyPhase() protocol.KeyPhaseBit { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "KeyPhase") + ret0, _ := ret[0].(protocol.KeyPhaseBit) + return ret0 +} + +// KeyPhase indicates an expected call of KeyPhase. +func (mr *MockShortHeaderSealerMockRecorder) KeyPhase() *MockShortHeaderSealerKeyPhaseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "KeyPhase", reflect.TypeOf((*MockShortHeaderSealer)(nil).KeyPhase)) + return &MockShortHeaderSealerKeyPhaseCall{Call: call} +} + +// MockShortHeaderSealerKeyPhaseCall wrap *gomock.Call +type MockShortHeaderSealerKeyPhaseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderSealerKeyPhaseCall) Return(arg0 protocol.KeyPhaseBit) *MockShortHeaderSealerKeyPhaseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderSealerKeyPhaseCall) Do(f func() protocol.KeyPhaseBit) *MockShortHeaderSealerKeyPhaseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderSealerKeyPhaseCall) DoAndReturn(f func() protocol.KeyPhaseBit) *MockShortHeaderSealerKeyPhaseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Overhead mocks base method. +func (m *MockShortHeaderSealer) Overhead() int { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Overhead") + ret0, _ := ret[0].(int) + return ret0 +} + +// Overhead indicates an expected call of Overhead. +func (mr *MockShortHeaderSealerMockRecorder) Overhead() *MockShortHeaderSealerOverheadCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Overhead", reflect.TypeOf((*MockShortHeaderSealer)(nil).Overhead)) + return &MockShortHeaderSealerOverheadCall{Call: call} +} + +// MockShortHeaderSealerOverheadCall wrap *gomock.Call +type MockShortHeaderSealerOverheadCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderSealerOverheadCall) Return(arg0 int) *MockShortHeaderSealerOverheadCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderSealerOverheadCall) Do(f func() int) *MockShortHeaderSealerOverheadCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderSealerOverheadCall) DoAndReturn(f func() int) *MockShortHeaderSealerOverheadCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Seal mocks base method. +func (m *MockShortHeaderSealer) Seal(dst, src []byte, packetNumber protocol.PacketNumber, associatedData []byte) []byte { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Seal", dst, src, packetNumber, associatedData) + ret0, _ := ret[0].([]byte) + return ret0 +} + +// Seal indicates an expected call of Seal. +func (mr *MockShortHeaderSealerMockRecorder) Seal(dst, src, packetNumber, associatedData any) *MockShortHeaderSealerSealCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Seal", reflect.TypeOf((*MockShortHeaderSealer)(nil).Seal), dst, src, packetNumber, associatedData) + return &MockShortHeaderSealerSealCall{Call: call} +} + +// MockShortHeaderSealerSealCall wrap *gomock.Call +type MockShortHeaderSealerSealCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockShortHeaderSealerSealCall) Return(arg0 []byte) *MockShortHeaderSealerSealCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockShortHeaderSealerSealCall) Do(f func([]byte, []byte, protocol.PacketNumber, []byte) []byte) *MockShortHeaderSealerSealCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockShortHeaderSealerSealCall) DoAndReturn(f func([]byte, []byte, protocol.PacketNumber, []byte) []byte) *MockShortHeaderSealerSealCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/internal/monotime/time.go b/third_party/quic-go/internal/monotime/time.go new file mode 100644 index 0000000..eda61dc --- /dev/null +++ b/third_party/quic-go/internal/monotime/time.go @@ -0,0 +1,90 @@ +// Package monotime provides a monotonic time representation that is useful for +// measuring elapsed time. +// It is designed as a memory optimized drop-in replacement for time.Time, with +// a monotime.Time consuming just 8 bytes instead of 24 bytes. +package monotime + +import ( + "time" +) + +// The absolute value doesn't matter, but it should be in the past, +// so that every timestamp obtained with Now() is non-zero, +// even on systems with low timer resolutions (e.g. Windows). +var start = time.Now().Add(-time.Hour) + +// A Time represents an instant in monotonic time. +// Times can be compared using the comparison operators, but the specific +// value is implementation-dependent and should not be relied upon. +// The zero value of Time doesn't have any specific meaning. +type Time int64 + +// Now returns the current monotonic time. +func Now() Time { + return Time(time.Since(start).Nanoseconds()) +} + +// Sub returns the duration t-t2. If the result exceeds the maximum (or minimum) +// value that can be stored in a Duration, the maximum (or minimum) duration +// will be returned. +// To compute t-d for a duration d, use t.Add(-d). +func (t Time) Sub(t2 Time) time.Duration { + return time.Duration(t - t2) +} + +// Add returns the time t+d. +func (t Time) Add(d time.Duration) Time { + return Time(int64(t) + d.Nanoseconds()) +} + +// After reports whether the time instant t is after t2. +func (t Time) After(t2 Time) bool { + return t > t2 +} + +// Before reports whether the time instant t is before t2. +func (t Time) Before(t2 Time) bool { + return t < t2 +} + +// IsZero reports whether t represents the zero time instant. +func (t Time) IsZero() bool { + return t == 0 +} + +// Equal reports whether t and t2 represent the same time instant. +func (t Time) Equal(t2 Time) bool { + return t == t2 +} + +// ToTime converts the monotonic time to a time.Time value. +// The returned time.Time will have the same instant as the monotonic time, +// but may be subject to clock adjustments. +func (t Time) ToTime() time.Time { + if t.IsZero() { + return time.Time{} + } + return start.Add(time.Duration(t)) +} + +// Since returns the time elapsed since t. It is shorthand for Now().Sub(t). +func Since(t Time) time.Duration { + return Now().Sub(t) +} + +// Until returns the duration until t. +// It is shorthand for t.Sub(Now()). +// If t is in the past, the returned duration will be negative. +func Until(t Time) time.Duration { + return time.Duration(t - Now()) +} + +// FromTime converts a time.Time to a monotonic Time. +// The conversion is relative to the package's start time and may lose +// precision if the time.Time is far from the start time. +func FromTime(t time.Time) Time { + if t.IsZero() { + return 0 + } + return Time(t.Sub(start).Nanoseconds()) +} diff --git a/third_party/quic-go/internal/monotime/time_test.go b/third_party/quic-go/internal/monotime/time_test.go new file mode 100644 index 0000000..390af6e --- /dev/null +++ b/third_party/quic-go/internal/monotime/time_test.go @@ -0,0 +1,78 @@ +package monotime + +import ( + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" +) + +func TestTimeRelations(t *testing.T) { + t1 := Now() + require.Equal(t, t1, t1) + require.False(t, t1.IsZero()) + + t2 := t1.Add(time.Second) + + require.False(t, t1.Equal(t2)) + require.False(t, t2.Equal(t1)) + + require.True(t, t2.After(t1)) + require.False(t, t1.After(t2)) + require.False(t, t2.Before(t1)) + + require.Equal(t, t2.Sub(t1), time.Second) + require.Equal(t, t1.Sub(t2), -time.Second) +} + +func TestSince(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + t1 := Now() + time.Sleep(time.Second) + require.Equal(t, Since(t1), time.Second) + require.Equal(t, Now().Sub(t1), time.Second) + time.Sleep(time.Minute) + require.Equal(t, Since(t1), time.Minute+time.Second) + require.Equal(t, Now().Sub(t1), time.Minute+time.Second) + }) +} + +func TestUntil(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + t1 := Now().Add(time.Minute) + require.Equal(t, Until(t1), time.Minute) + require.Equal(t, t1.Sub(Now()), time.Minute) + time.Sleep(15 * time.Second) + require.Equal(t, Until(t1), 45*time.Second) + require.Equal(t, t1.Sub(Now()), 45*time.Second) + }) +} + +func TestConversions(t *testing.T) { + t1 := Now() + t1Time := t1.ToTime() + require.Equal(t, FromTime(t1Time), t1) + require.Zero(t, t1Time.Sub(t1.ToTime())) + + var zeroTime time.Time + require.Zero(t, FromTime(zeroTime)) + require.Zero(t, FromTime(zeroTime)) + + var zero Time + require.True(t, zero.ToTime().IsZero()) +} + +func BenchmarkNow(b *testing.B) { + b.Run("Now", func(b *testing.B) { + for b.Loop() { + _ = Now() + } + }) + + b.Run("time.Now", func(b *testing.B) { + for b.Loop() { + _ = time.Now() + } + }) +} diff --git a/third_party/quic-go/internal/protocol/connection_id.go b/third_party/quic-go/internal/protocol/connection_id.go new file mode 100644 index 0000000..5c59fc3 --- /dev/null +++ b/third_party/quic-go/internal/protocol/connection_id.go @@ -0,0 +1,127 @@ +package protocol + +import ( + "crypto/rand" + "encoding/hex" + "errors" + "io" +) + +var ErrInvalidConnectionIDLen = errors.New("invalid Connection ID length") + +// An ArbitraryLenConnectionID is a QUIC Connection ID able to represent Connection IDs according to RFC 8999. +// Future QUIC versions might allow connection ID lengths up to 255 bytes, while QUIC v1 +// restricts the length to 20 bytes. +type ArbitraryLenConnectionID []byte + +func (c ArbitraryLenConnectionID) Len() int { + return len(c) +} + +func (c ArbitraryLenConnectionID) Bytes() []byte { + return c +} + +func (c ArbitraryLenConnectionID) String() string { + if c.Len() == 0 { + return "(empty)" + } + return hex.EncodeToString(c.Bytes()) +} + +const maxConnectionIDLen = 20 + +// A ConnectionID in QUIC +type ConnectionID struct { + b [20]byte + l uint8 +} + +// GenerateConnectionID generates a connection ID using cryptographic random +func GenerateConnectionID(l int) (ConnectionID, error) { + var c ConnectionID + c.l = uint8(l) + _, err := rand.Read(c.b[:l]) + return c, err +} + +// ParseConnectionID interprets b as a Connection ID. +// It panics if b is longer than 20 bytes. +func ParseConnectionID(b []byte) ConnectionID { + if len(b) > maxConnectionIDLen { + panic("invalid conn id length") + } + var c ConnectionID + c.l = uint8(len(b)) + copy(c.b[:c.l], b) + return c +} + +// GenerateConnectionIDForInitial generates a connection ID for the Initial packet. +// It uses a length randomly chosen between 8 and 20 bytes. +func GenerateConnectionIDForInitial() (ConnectionID, error) { + r := make([]byte, 1) + if _, err := rand.Read(r); err != nil { + return ConnectionID{}, err + } + l := MinConnectionIDLenInitial + int(r[0])%(maxConnectionIDLen-MinConnectionIDLenInitial+1) + return GenerateConnectionID(l) +} + +// ChromeConnectionIDLenInitial is the initial destination connection ID length a +// Chrome-parroting client uses. quic-go randomizes the length instead, which is +// visible on the wire. +const ChromeConnectionIDLenInitial = 8 + +// GenerateChromeConnectionIDForInitial generates an initial destination +// connection ID of the fixed length above. +func GenerateChromeConnectionIDForInitial() (ConnectionID, error) { + return GenerateConnectionID(ChromeConnectionIDLenInitial) +} + +// ReadConnectionID reads a connection ID of length len from the given io.Reader. +// It returns io.EOF if there are not enough bytes to read. +func ReadConnectionID(r io.Reader, l int) (ConnectionID, error) { + var c ConnectionID + if l == 0 { + return c, nil + } + if l > maxConnectionIDLen { + return c, ErrInvalidConnectionIDLen + } + c.l = uint8(l) + _, err := io.ReadFull(r, c.b[:l]) + if err == io.ErrUnexpectedEOF { + return c, io.EOF + } + return c, err +} + +// Len returns the length of the connection ID in bytes +func (c ConnectionID) Len() int { + return int(c.l) +} + +// Bytes returns the byte representation +func (c ConnectionID) Bytes() []byte { + return c.b[:c.l] +} + +func (c ConnectionID) String() string { + if c.Len() == 0 { + return "(empty)" + } + return hex.EncodeToString(c.Bytes()) +} + +type DefaultConnectionIDGenerator struct { + ConnLen int +} + +func (d *DefaultConnectionIDGenerator) GenerateConnectionID() (ConnectionID, error) { + return GenerateConnectionID(d.ConnLen) +} + +func (d *DefaultConnectionIDGenerator) ConnectionIDLen() int { + return d.ConnLen +} diff --git a/third_party/quic-go/internal/protocol/connection_id_test.go b/third_party/quic-go/internal/protocol/connection_id_test.go new file mode 100644 index 0000000..f658ca2 --- /dev/null +++ b/third_party/quic-go/internal/protocol/connection_id_test.go @@ -0,0 +1,92 @@ +package protocol + +import ( + "bytes" + "crypto/rand" + "io" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGenerateRandomConnectionIDs(t *testing.T) { + c1, err := GenerateConnectionID(8) + require.NoError(t, err) + require.NotZero(t, c1) + require.Equal(t, 8, c1.Len()) + c2, err := GenerateConnectionID(8) + require.NoError(t, err) + require.NotEqual(t, c1, c2) + require.Equal(t, 8, c2.Len()) +} + +func TestGenerateRandomLengthDestinationConnectionIDs(t *testing.T) { + var has8ByteConnID, has20ByteConnID bool + for range 1000 { + c, err := GenerateConnectionIDForInitial() + require.NoError(t, err) + require.GreaterOrEqual(t, c.Len(), 8) + require.LessOrEqual(t, c.Len(), 20) + if c.Len() == 8 { + has8ByteConnID = true + } + if c.Len() == 20 { + has20ByteConnID = true + } + } + require.True(t, has8ByteConnID) + require.True(t, has20ByteConnID) +} + +func TestConnectionID(t *testing.T) { + buf := bytes.NewBuffer([]byte{0xde, 0xad, 0xbe, 0xef, 0x42}) + c, err := ReadConnectionID(buf, 5) + require.NoError(t, err) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef, 0x42}, c.Bytes()) + require.Equal(t, 5, c.Len()) + require.Equal(t, "deadbeef42", c.String()) + + // too few bytes + _, err = ReadConnectionID(buf, 10) + require.Equal(t, io.EOF, err) + + // zero length + c2, err := ReadConnectionID(buf, 0) + require.NoError(t, err) + require.Zero(t, c2.Len()) + + // connection ID can have a length of a maximum of 20 bytes + buf2 := bytes.NewBuffer(make([]byte, 21)) + _, err = ReadConnectionID(buf2, 21) + require.Equal(t, ErrInvalidConnectionIDLen, err) +} + +func TestConnectionIDZeroValue(t *testing.T) { + var c ConnectionID + require.Zero(t, c.Len()) + require.Empty(t, c.Bytes()) + require.Equal(t, "(empty)", (ConnectionID{}).String()) +} + +// The string representation of a connection ID is used in qlog, so it should be fast. +func BenchmarkConnectionIDStringer(b *testing.B) { + c := ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0x42}) + b.ReportAllocs() + for b.Loop() { + _ = c.String() + } +} + +func TestArbitraryLenConnectionID(t *testing.T) { + b := make([]byte, 42) + rand.Read(b) + c := ArbitraryLenConnectionID(b) + require.Equal(t, b, c.Bytes()) + require.Equal(t, 42, c.Len()) +} + +func TestArbitraryLenConnectionIDStringer(t *testing.T) { + require.Equal(t, "(empty)", (ArbitraryLenConnectionID{}).String()) + c := ArbitraryLenConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0x42}) + require.Equal(t, "deadbeef42", c.String()) +} diff --git a/third_party/quic-go/internal/protocol/encryption_level.go b/third_party/quic-go/internal/protocol/encryption_level.go new file mode 100644 index 0000000..40aa331 --- /dev/null +++ b/third_party/quic-go/internal/protocol/encryption_level.go @@ -0,0 +1,65 @@ +package protocol + +import ( + "crypto/tls" + "fmt" +) + +// EncryptionLevel is the encryption level +// Default value is Unencrypted +type EncryptionLevel uint8 + +const ( + // EncryptionInitial is the Initial encryption level + EncryptionInitial EncryptionLevel = 1 + iota + // EncryptionHandshake is the Handshake encryption level + EncryptionHandshake + // Encryption0RTT is the 0-RTT encryption level + Encryption0RTT + // Encryption1RTT is the 1-RTT encryption level + Encryption1RTT +) + +func (e EncryptionLevel) String() string { + switch e { + case EncryptionInitial: + return "Initial" + case EncryptionHandshake: + return "Handshake" + case Encryption0RTT: + return "0-RTT" + case Encryption1RTT: + return "1-RTT" + } + return "unknown" +} + +func (e EncryptionLevel) ToTLSEncryptionLevel() tls.QUICEncryptionLevel { + switch e { + case EncryptionInitial: + return tls.QUICEncryptionLevelInitial + case EncryptionHandshake: + return tls.QUICEncryptionLevelHandshake + case Encryption1RTT: + return tls.QUICEncryptionLevelApplication + case Encryption0RTT: + return tls.QUICEncryptionLevelEarly + default: + panic(fmt.Sprintf("unexpected encryption level: %s", e)) + } +} + +func FromTLSEncryptionLevel(e tls.QUICEncryptionLevel) EncryptionLevel { + switch e { + case tls.QUICEncryptionLevelInitial: + return EncryptionInitial + case tls.QUICEncryptionLevelHandshake: + return EncryptionHandshake + case tls.QUICEncryptionLevelApplication: + return Encryption1RTT + case tls.QUICEncryptionLevelEarly: + return Encryption0RTT + default: + panic(fmt.Sprintf("unexpect encryption level: %s", e)) + } +} diff --git a/third_party/quic-go/internal/protocol/encryption_level_test.go b/third_party/quic-go/internal/protocol/encryption_level_test.go new file mode 100644 index 0000000..7792a36 --- /dev/null +++ b/third_party/quic-go/internal/protocol/encryption_level_test.go @@ -0,0 +1,40 @@ +package protocol + +import ( + "crypto/tls" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestEncryptionLevelNonZeroValue(t *testing.T) { + require.NotZero(t, EncryptionInitial*EncryptionHandshake*Encryption0RTT*Encryption1RTT) +} + +func TestEncryptionLevelConversion(t *testing.T) { + testCases := []struct { + quicLevel EncryptionLevel + tlsLevel tls.QUICEncryptionLevel + }{ + {EncryptionInitial, tls.QUICEncryptionLevelInitial}, + {EncryptionHandshake, tls.QUICEncryptionLevelHandshake}, + {Encryption1RTT, tls.QUICEncryptionLevelApplication}, + {Encryption0RTT, tls.QUICEncryptionLevelEarly}, + } + + for _, tc := range testCases { + t.Run(tc.quicLevel.String(), func(t *testing.T) { + // conversion from QUIC to TLS encryption level + require.Equal(t, tc.tlsLevel, tc.quicLevel.ToTLSEncryptionLevel()) + // conversion from TLS to QUIC encryption level + require.Equal(t, tc.quicLevel, FromTLSEncryptionLevel(tc.tlsLevel)) + }) + } +} + +func TestEncryptionLevelStringRepresentation(t *testing.T) { + require.Equal(t, "Initial", EncryptionInitial.String()) + require.Equal(t, "Handshake", EncryptionHandshake.String()) + require.Equal(t, "0-RTT", Encryption0RTT.String()) + require.Equal(t, "1-RTT", Encryption1RTT.String()) +} diff --git a/third_party/quic-go/internal/protocol/key_phase.go b/third_party/quic-go/internal/protocol/key_phase.go new file mode 100644 index 0000000..edd740c --- /dev/null +++ b/third_party/quic-go/internal/protocol/key_phase.go @@ -0,0 +1,36 @@ +package protocol + +// KeyPhase is the key phase +type KeyPhase uint64 + +// Bit determines the key phase bit +func (p KeyPhase) Bit() KeyPhaseBit { + if p%2 == 0 { + return KeyPhaseZero + } + return KeyPhaseOne +} + +// KeyPhaseBit is the key phase bit +type KeyPhaseBit uint8 + +const ( + // KeyPhaseUndefined is an undefined key phase + KeyPhaseUndefined KeyPhaseBit = iota + // KeyPhaseZero is key phase 0 + KeyPhaseZero + // KeyPhaseOne is key phase 1 + KeyPhaseOne +) + +func (p KeyPhaseBit) String() string { + //nolint:exhaustive + switch p { + case KeyPhaseZero: + return "0" + case KeyPhaseOne: + return "1" + default: + return "undefined" + } +} diff --git a/third_party/quic-go/internal/protocol/key_phase_test.go b/third_party/quic-go/internal/protocol/key_phase_test.go new file mode 100644 index 0000000..34b3351 --- /dev/null +++ b/third_party/quic-go/internal/protocol/key_phase_test.go @@ -0,0 +1,26 @@ +package protocol + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestKeyPhaseBitDefaultValue(t *testing.T) { + var k KeyPhaseBit + require.Equal(t, KeyPhaseUndefined, k) +} + +func TestKeyPhaseStringRepresentation(t *testing.T) { + require.Equal(t, "0", KeyPhaseZero.String()) + require.Equal(t, "1", KeyPhaseOne.String()) +} + +func TestKeyPhaseToBit(t *testing.T) { + require.Equal(t, KeyPhaseZero, KeyPhase(0).Bit()) + require.Equal(t, KeyPhaseZero, KeyPhase(2).Bit()) + require.Equal(t, KeyPhaseZero, KeyPhase(4).Bit()) + require.Equal(t, KeyPhaseOne, KeyPhase(1).Bit()) + require.Equal(t, KeyPhaseOne, KeyPhase(3).Bit()) + require.Equal(t, KeyPhaseOne, KeyPhase(5).Bit()) +} diff --git a/third_party/quic-go/internal/protocol/packet_number.go b/third_party/quic-go/internal/protocol/packet_number.go new file mode 100644 index 0000000..f249a8e --- /dev/null +++ b/third_party/quic-go/internal/protocol/packet_number.go @@ -0,0 +1,84 @@ +package protocol + +// A PacketNumber in QUIC +type PacketNumber int64 + +// InvalidPacketNumber is a packet number that is never sent. +// In QUIC, 0 is a valid packet number. +const InvalidPacketNumber PacketNumber = -1 + +// PacketNumberLen is the length of the packet number in bytes +type PacketNumberLen uint8 + +const ( + // PacketNumberLen1 is a packet number length of 1 byte + PacketNumberLen1 PacketNumberLen = 1 + // PacketNumberLen2 is a packet number length of 2 bytes + PacketNumberLen2 PacketNumberLen = 2 + // PacketNumberLen3 is a packet number length of 3 bytes + PacketNumberLen3 PacketNumberLen = 3 + // PacketNumberLen4 is a packet number length of 4 bytes + PacketNumberLen4 PacketNumberLen = 4 +) + +// DecodePacketNumber calculates the packet number based its length and the last seen packet number +// This function is taken from https://www.rfc-editor.org/rfc/rfc9000.html#section-a.3. +func DecodePacketNumber(length PacketNumberLen, largest PacketNumber, truncated PacketNumber) PacketNumber { + expected := largest + 1 + win := PacketNumber(1 << (length * 8)) + hwin := win / 2 + mask := win - 1 + candidate := (expected & ^mask) | truncated + if candidate <= expected-hwin && candidate < 1<<62-win { + return candidate + win + } + if candidate > expected+hwin && candidate >= win { + return candidate - win + } + return candidate +} + +// PacketNumberLengthForHeader gets the length of the packet number for the public header +// it never chooses a PacketNumberLen of 1 byte, since this is too short under certain circumstances +func PacketNumberLengthForHeader(pn, largestAcked PacketNumber) PacketNumberLen { + var numUnacked PacketNumber + if largestAcked == InvalidPacketNumber { + numUnacked = pn + 1 + } else { + numUnacked = pn - largestAcked + } + if numUnacked < 1<<(16-1) { + return PacketNumberLen2 + } + if numUnacked < 1<<(24-1) { + return PacketNumberLen3 + } + return PacketNumberLen4 +} + +// PacketNumberLengthForHeaderChrome sizes the packet number the way the imitated +// client does: take the larger of the unacknowledged range and the congestion +// window measured in packets, quadruple it, and use the shortest length that can +// encode that. Both this and PacketNumberLengthForHeader conform to RFC 9000 +// section 17.1, but they differ on the wire in two visible ways. +// +// The congestion window sets a floor, so a single byte is used for far longer +// than a rule based on the unacknowledged range alone would allow: the window +// only has to stay under a quarter of the one-byte range. Three-byte packet +// numbers are also never produced, and quic-go's default does produce them. +func PacketNumberLengthForHeaderChrome(pn, largestAcked PacketNumber, cwndPackets PacketNumber) PacketNumberLen { + var numUnacked PacketNumber + if largestAcked == InvalidPacketNumber { + numUnacked = pn + 1 + } else { + numUnacked = pn - largestAcked + } + delta := max(numUnacked, cwndPackets) + if delta < 1<<8/4 { + return PacketNumberLen1 + } + if delta < 1<<16/4 { + return PacketNumberLen2 + } + return PacketNumberLen4 +} diff --git a/third_party/quic-go/internal/protocol/packet_number_test.go b/third_party/quic-go/internal/protocol/packet_number_test.go new file mode 100644 index 0000000..46b1901 --- /dev/null +++ b/third_party/quic-go/internal/protocol/packet_number_test.go @@ -0,0 +1,81 @@ +package protocol + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestInvalidPacketNumberIsSmallerThanAllValidPacketNumbers(t *testing.T) { + require.Less(t, InvalidPacketNumber, PacketNumber(0)) +} + +func TestPacketNumberLenHasCorrectValue(t *testing.T) { + require.EqualValues(t, 1, PacketNumberLen1) + require.EqualValues(t, 2, PacketNumberLen2) + require.EqualValues(t, 3, PacketNumberLen3) + require.EqualValues(t, 4, PacketNumberLen4) +} + +func TestDecodePacketNumber(t *testing.T) { + require.Equal(t, PacketNumber(255), DecodePacketNumber(PacketNumberLen1, 10, 255)) + require.Equal(t, PacketNumber(0), DecodePacketNumber(PacketNumberLen1, 10, 0)) + require.Equal(t, PacketNumber(256), DecodePacketNumber(PacketNumberLen1, 127, 0)) + require.Equal(t, PacketNumber(256), DecodePacketNumber(PacketNumberLen1, 128, 0)) + require.Equal(t, PacketNumber(256), DecodePacketNumber(PacketNumberLen1, 256+126, 0)) + require.Equal(t, PacketNumber(512), DecodePacketNumber(PacketNumberLen1, 256+127, 0)) + require.Equal(t, PacketNumber(0xffff), DecodePacketNumber(PacketNumberLen2, 0xffff, 0xffff)) + require.Equal(t, PacketNumber(0xffff), DecodePacketNumber(PacketNumberLen2, 0xffff+1, 0xffff)) + + // example from https://www.rfc-editor.org/rfc/rfc9000.html#section-a.3 + require.Equal(t, PacketNumber(0xa82f9b32), DecodePacketNumber(PacketNumberLen2, 0xa82f30ea, 0x9b32)) +} + +func TestPacketNumberLengthForHeader(t *testing.T) { + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeader(1, InvalidPacketNumber)) + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeader(1<<15-2, InvalidPacketNumber)) + require.Equal(t, PacketNumberLen3, PacketNumberLengthForHeader(1<<15-1, InvalidPacketNumber)) + require.Equal(t, PacketNumberLen3, PacketNumberLengthForHeader(1<<23-2, InvalidPacketNumber)) + require.Equal(t, PacketNumberLen4, PacketNumberLengthForHeader(1<<23-1, InvalidPacketNumber)) + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeader(1<<15+9, 10)) + require.Equal(t, PacketNumberLen3, PacketNumberLengthForHeader(1<<15+10, 10)) + require.Equal(t, PacketNumberLen3, PacketNumberLengthForHeader(1<<23+99, 100)) + require.Equal(t, PacketNumberLen4, PacketNumberLengthForHeader(1<<23+100, 100)) + // examples from https://www.rfc-editor.org/rfc/rfc9000.html#section-a.2 + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeader(0xac5c02, 0xabe8b3)) + require.Equal(t, PacketNumberLen3, PacketNumberLengthForHeader(0xace8fe, 0xabe8b3)) +} + +func TestPacketNumberLengthForHeaderChrome(t *testing.T) { + // quic-go's floor is 2 bytes; the parroted client uses 1 wherever the rule + // allows, so every packet would otherwise be one byte longer. + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeader(1, InvalidPacketNumber)) + require.Equal(t, PacketNumberLen1, PacketNumberLengthForHeaderChrome(1, InvalidPacketNumber, 0)) + + // The threshold is on four times the larger of the unacked range and the + // congestion window in packets. Pin both sides of each boundary. + require.Equal(t, PacketNumberLen1, PacketNumberLengthForHeaderChrome(63, 0, 0)) + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeaderChrome(64, 0, 0)) + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeaderChrome(1<<14-1, 0, 0)) + require.Equal(t, PacketNumberLen4, PacketNumberLengthForHeaderChrome(1<<14, 0, 0)) + + // The congestion window is a floor, so a wide-open window forces a longer + // packet number even when everything is acknowledged. + require.Equal(t, PacketNumberLen1, PacketNumberLengthForHeaderChrome(1000, 999, 63)) + require.Equal(t, PacketNumberLen2, PacketNumberLengthForHeaderChrome(1000, 999, 64)) + + // The initial congestion window must still leave room for 1 byte, otherwise + // the whole handshake is a byte wider than the client being imitated. + require.Equal(t, PacketNumberLen1, PacketNumberLengthForHeaderChrome(1, InvalidPacketNumber, 32)) + + // Three-byte packet numbers are never produced, though the default does + // produce them; a lone 3-byte packet number would be a giveaway. + require.Equal(t, PacketNumberLen3, PacketNumberLengthForHeader(1<<15+10, 10)) + for _, cwnd := range []PacketNumber{0, 32, 1000, 1 << 20} { + for _, pn := range []PacketNumber{1, 63, 64, 1 << 10, 1<<14 - 1, 1 << 14, 1 << 20, 1 << 25} { + require.NotEqual(t, PacketNumberLen3, + PacketNumberLengthForHeaderChrome(pn, InvalidPacketNumber, cwnd), + "pn=%d cwnd=%d", pn, cwnd) + } + } +} diff --git a/third_party/quic-go/internal/protocol/params.go b/third_party/quic-go/internal/protocol/params.go new file mode 100644 index 0000000..93297b6 --- /dev/null +++ b/third_party/quic-go/internal/protocol/params.go @@ -0,0 +1,169 @@ +package protocol + +import "time" + +// DesiredReceiveBufferSize is the kernel UDP receive buffer size that we'd like to use. +const DesiredReceiveBufferSize = (1 << 20) * 8 // 8 MB + +// DesiredSendBufferSize is the kernel UDP send buffer size that we'd like to use. +const DesiredSendBufferSize = (1 << 20) * 8 // 8 MB + +// InitialPacketSize is the initial (before Path MTU discovery) maximum packet size used. +const InitialPacketSize = 1280 + +// MaxCongestionWindowPackets is the maximum congestion window in packet. +const MaxCongestionWindowPackets = 20000 + +// MaxUndecryptablePackets limits the number of undecryptable packets that are queued in the connection. +const MaxUndecryptablePackets = 32 + +// ConnectionFlowControlMultiplier determines how much larger the connection flow control windows needs to be relative to any stream's flow control window +// This is the value that Chromium is using +const ConnectionFlowControlMultiplier = 1.5 + +// DefaultInitialMaxStreamData is the default initial stream-level flow control window for receiving data +const DefaultInitialMaxStreamData = (1 << 20) * 2 // 2 MB + +// DefaultInitialMaxData is the connection-level flow control window for receiving data +const DefaultInitialMaxData = ConnectionFlowControlMultiplier * DefaultInitialMaxStreamData + +// DefaultMaxReceiveStreamFlowControlWindow is the default maximum stream-level flow control window for receiving data +const DefaultMaxReceiveStreamFlowControlWindow = 6 * (1 << 20) // 6 MB + +// DefaultMaxReceiveConnectionFlowControlWindow is the default connection-level flow control window for receiving data +const DefaultMaxReceiveConnectionFlowControlWindow = 15 * (1 << 20) // 15 MB + +// WindowUpdateThreshold is the fraction of the receive window that has to be consumed before an higher offset is advertised to the client +const WindowUpdateThreshold = 0.25 + +// DefaultMaxIncomingStreams is the maximum number of streams that a peer may open +const DefaultMaxIncomingStreams = 100 + +// DefaultMaxIncomingUniStreams is the maximum number of unidirectional streams that a peer may open +const DefaultMaxIncomingUniStreams = 100 + +// MaxServerUnprocessedPackets is the max number of packets stored in the server that are not yet processed. +const MaxServerUnprocessedPackets = 1024 + +// MaxConnUnprocessedPackets is the max number of packets stored in each connection that are not yet processed. +const MaxConnUnprocessedPackets = 256 + +// SkipPacketInitialPeriod is the initial period length used for packet number skipping to prevent an Optimistic ACK attack. +// Every time a packet number is skipped, the period is doubled, up to SkipPacketMaxPeriod. +const SkipPacketInitialPeriod PacketNumber = 256 + +// SkipPacketMaxPeriod is the maximum period length used for packet number skipping. +const SkipPacketMaxPeriod PacketNumber = 128 * 1024 + +// MaxAcceptQueueSize is the maximum number of connections that the server queues for accepting. +// If the queue is full, new connection attempts will be rejected. +const MaxAcceptQueueSize = 32 + +// TokenValidity is the duration that a (non-retry) token is considered valid +const TokenValidity = 24 * time.Hour + +// MaxOutstandingSentPackets is maximum number of packets saved for retransmission. +// When reached, it imposes a soft limit on sending new packets: +// Sending ACKs and retransmission is still allowed, but now new regular packets can be sent. +const MaxOutstandingSentPackets = 2 * MaxCongestionWindowPackets + +// MaxTrackedSentPackets is maximum number of sent packets saved for retransmission. +// When reached, no more packets will be sent. +// This value *must* be larger than MaxOutstandingSentPackets. +const MaxTrackedSentPackets = MaxOutstandingSentPackets * 5 / 4 + +// MaxNonAckElicitingAcks is the maximum number of packets containing an ACK, +// but no ack-eliciting frames, that we send in a row +const MaxNonAckElicitingAcks = 19 + +// MaxStreamFrameSorterGaps is the maximum number of gaps between received StreamFrames +// prevents DoS attacks against the streamFrameSorter +const MaxStreamFrameSorterGaps = 20000 + +// MinStreamFrameBufferSize is the minimum data length of a received STREAM frame +// that we use the buffer for. This protects against a DoS where an attacker would send us +// very small STREAM frames to consume a lot of memory. +const MinStreamFrameBufferSize = 128 + +// MinCoalescedPacketSize is the minimum size of a coalesced packet that we pack. +// If a packet has less than this number of bytes, we won't coalesce any more packets onto it. +const MinCoalescedPacketSize = 128 + +// MaxCryptoStreamOffset is the maximum offset allowed on any of the crypto streams. +// This limits the size of the ClientHello and Certificates that can be received. +const MaxCryptoStreamOffset = 16 * (1 << 10) + +// MinRemoteIdleTimeout is the minimum value that we accept for the remote idle timeout +const MinRemoteIdleTimeout = 5 * time.Second + +// DefaultIdleTimeout is the default idle timeout +const DefaultIdleTimeout = 30 * time.Second + +// DefaultHandshakeIdleTimeout is the default idle timeout used before handshake completion. +const DefaultHandshakeIdleTimeout = 5 * time.Second + +// MinStreamFrameSize is the minimum size that has to be left in a packet, so that we add another STREAM frame. +// This avoids splitting up STREAM frames into small pieces, which has 2 advantages: +// 1. it reduces the framing overhead +// 2. it reduces the head-of-line blocking, when a packet is lost +const MinStreamFrameSize ByteCount = 128 + +// MaxPostHandshakeCryptoFrameSize is the maximum size of CRYPTO frames +// we send after the handshake completes. +const MaxPostHandshakeCryptoFrameSize = 1000 + +// MaxNumAckRanges is the maximum number of ACK ranges that we send in an ACK frame. +// It also serves as a limit for the packet history. +// If at any point we keep track of more ranges, old ranges are discarded. +// +// This value also guarantees that ACK Range Count value in the ACK frame can be encoded +// in a single byte varint. +const MaxNumAckRanges = 64 + +// MinPacingDelay is the minimum duration that is used for packet pacing +// If the packet packing frequency is higher, multiple packets might be sent at once. +// Example: For a packet pacing delay of 200μs, we would send 5 packets at once, wait for 1ms, and so forth. +const MinPacingDelay = time.Millisecond + +// DefaultConnectionIDLength is the connection ID length that is used for multiplexed connections +// if no other value is configured. +const DefaultConnectionIDLength = 4 + +// MaxActiveConnectionIDs is the number of connection IDs that we're storing. +const MaxActiveConnectionIDs = 4 + +// MaxIssuedConnectionIDs is the maximum number of connection IDs that we're issuing at the same time. +const MaxIssuedConnectionIDs = 6 + +// PacketsPerConnectionID is the number of packets we send using one connection ID. +// If the peer provices us with enough new connection IDs, we switch to a new connection ID. +const PacketsPerConnectionID = 10000 + +// AckDelayExponent is the ack delay exponent used when sending ACKs. +const AckDelayExponent = 3 + +// Estimated timer granularity. +// The loss detection timer will not be set to a value smaller than granularity. +const TimerGranularity = time.Millisecond + +// MaxAckDelay is the maximum time by which we delay sending ACKs. +const MaxAckDelay = 25 * time.Millisecond + +// MaxAckDelayInclGranularity is the max_ack_delay including the timer granularity. +// This is the value that should be advertised to the peer. +const MaxAckDelayInclGranularity = MaxAckDelay + TimerGranularity + +// KeyUpdateInterval is the maximum number of packets we send or receive before initiating a key update. +const KeyUpdateInterval = 100 * 1000 + +// Max0RTTQueueingDuration is the maximum time that we store 0-RTT packets in order to wait for the corresponding Initial to be received. +const Max0RTTQueueingDuration = 100 * time.Millisecond + +// Max0RTTQueues is the maximum number of connections that we buffer 0-RTT packets for. +const Max0RTTQueues = 32 + +// Max0RTTQueueLen is the maximum number of 0-RTT packets that we buffer for each connection. +// When a new connection is created, all buffered packets are passed to the connection immediately. +// To avoid blocking, this value has to be smaller than MaxConnUnprocessedPackets. +// To avoid packets being dropped as undecryptable by the connection, this value has to be smaller than MaxUndecryptablePackets. +const Max0RTTQueueLen = 31 diff --git a/third_party/quic-go/internal/protocol/params_test.go b/third_party/quic-go/internal/protocol/params_test.go new file mode 100644 index 0000000..48e023d --- /dev/null +++ b/third_party/quic-go/internal/protocol/params_test.go @@ -0,0 +1,13 @@ +package protocol + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestPacketQueueCapacities(t *testing.T) { + // Ensure that the session can queue more packets than the 0-RTT queue + require.Greater(t, MaxConnUnprocessedPackets, Max0RTTQueueLen) + require.Greater(t, MaxUndecryptablePackets, Max0RTTQueueLen) +} diff --git a/third_party/quic-go/internal/protocol/perspective.go b/third_party/quic-go/internal/protocol/perspective.go new file mode 100644 index 0000000..5a29d3c --- /dev/null +++ b/third_party/quic-go/internal/protocol/perspective.go @@ -0,0 +1,26 @@ +package protocol + +// Perspective determines if we're acting as a server or a client +type Perspective int + +// the perspectives +const ( + PerspectiveServer Perspective = 1 + PerspectiveClient Perspective = 2 +) + +// Opposite returns the perspective of the peer +func (p Perspective) Opposite() Perspective { + return 3 - p +} + +func (p Perspective) String() string { + switch p { + case PerspectiveServer: + return "server" + case PerspectiveClient: + return "client" + default: + return "invalid perspective" + } +} diff --git a/third_party/quic-go/internal/protocol/perspective_test.go b/third_party/quic-go/internal/protocol/perspective_test.go new file mode 100644 index 0000000..121b478 --- /dev/null +++ b/third_party/quic-go/internal/protocol/perspective_test.go @@ -0,0 +1,18 @@ +package protocol + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestPerspectiveOpposite(t *testing.T) { + require.Equal(t, PerspectiveServer, PerspectiveClient.Opposite()) + require.Equal(t, PerspectiveClient, PerspectiveServer.Opposite()) +} + +func TestPerspectiveStringer(t *testing.T) { + require.Equal(t, "client", PerspectiveClient.String()) + require.Equal(t, "server", PerspectiveServer.String()) + require.Equal(t, "invalid perspective", Perspective(0).String()) +} diff --git a/third_party/quic-go/internal/protocol/protocol.go b/third_party/quic-go/internal/protocol/protocol.go new file mode 100644 index 0000000..b1109ff --- /dev/null +++ b/third_party/quic-go/internal/protocol/protocol.go @@ -0,0 +1,156 @@ +package protocol + +import ( + "fmt" + "time" +) + +// The PacketType is the Long Header Type +type PacketType uint8 + +const ( + // PacketTypeInitial is the packet type of an Initial packet + PacketTypeInitial PacketType = 1 + iota + // PacketTypeRetry is the packet type of a Retry packet + PacketTypeRetry + // PacketTypeHandshake is the packet type of a Handshake packet + PacketTypeHandshake + // PacketType0RTT is the packet type of a 0-RTT packet + PacketType0RTT +) + +func (t PacketType) String() string { + switch t { + case PacketTypeInitial: + return "Initial" + case PacketTypeRetry: + return "Retry" + case PacketTypeHandshake: + return "Handshake" + case PacketType0RTT: + return "0-RTT Protected" + default: + return fmt.Sprintf("unknown packet type: %d", t) + } +} + +type ECN uint8 + +const ( + ECNUnsupported ECN = iota + ECNNon // 00 + ECT1 // 01 + ECT0 // 10 + ECNCE // 11 +) + +func ParseECNHeaderBits(bits byte) ECN { + switch bits { + case 0: + return ECNNon + case 0b00000010: + return ECT0 + case 0b00000001: + return ECT1 + case 0b00000011: + return ECNCE + default: + panic("invalid ECN bits") + } +} + +func (e ECN) ToHeaderBits() byte { + //nolint:exhaustive // There are only 4 values. + switch e { + case ECNNon: + return 0 + case ECT0: + return 0b00000010 + case ECT1: + return 0b00000001 + case ECNCE: + return 0b00000011 + default: + panic("ECN unsupported") + } +} + +func (e ECN) String() string { + switch e { + case ECNUnsupported: + return "ECN unsupported" + case ECNNon: + return "Not-ECT" + case ECT1: + return "ECT(1)" + case ECT0: + return "ECT(0)" + case ECNCE: + return "CE" + default: + return fmt.Sprintf("invalid ECN value: %d", e) + } +} + +// A ByteCount in QUIC +type ByteCount int64 + +// MaxByteCount is the maximum value of a ByteCount +const MaxByteCount = ByteCount(1<<62 - 1) + +// InvalidByteCount is an invalid byte count +const InvalidByteCount ByteCount = -1 + +// A StatelessResetToken is a stateless reset token. +type StatelessResetToken [16]byte + +// MaxPacketBufferSize maximum packet size of any QUIC packet, based on +// ethernet's max size, minus the IP and UDP headers. IPv6 has a 40 byte header, +// UDP adds an additional 8 bytes. This is a total overhead of 48 bytes. +// Ethernet's max packet size is 1500 bytes, 1500 - 48 = 1452. +const MaxPacketBufferSize = 1452 + +// MaxLargePacketBufferSize is used when using GSO +const MaxLargePacketBufferSize = 20 * 1024 + +// MinInitialPacketSize is the minimum size an Initial packet is required to have. +const MinInitialPacketSize = 1200 + +// MinUnknownVersionPacketSize is the minimum size a packet with an unknown version +// needs to have in order to trigger a Version Negotiation packet. +const MinUnknownVersionPacketSize = MinInitialPacketSize + +// MinStatelessResetSize is the minimum size of a stateless reset packet that we send +const MinStatelessResetSize = 1 /* first byte */ + 20 /* max. conn ID length */ + 4 /* max. packet number length */ + 1 /* min. payload length */ + 16 /* token */ + +// MinReceivedStatelessResetSize is the minimum size of a received stateless reset, +// as specified in section 10.3 of RFC 9000. +const MinReceivedStatelessResetSize = 5 + 16 + +// MinConnectionIDLenInitial is the minimum length of the destination connection ID on an Initial packet. +const MinConnectionIDLenInitial = 8 + +// DefaultAckDelayExponent is the default ack delay exponent +const DefaultAckDelayExponent = 3 + +// DefaultActiveConnectionIDLimit is the default active connection ID limit +const DefaultActiveConnectionIDLimit = 2 + +// MaxAckDelayExponent is the maximum ack delay exponent +const MaxAckDelayExponent = 20 + +// DefaultMaxAckDelay is the default max_ack_delay +const DefaultMaxAckDelay = 25 * time.Millisecond + +// MaxMaxAckDelay is the maximum max_ack_delay +const MaxMaxAckDelay = (1<<14 - 1) * time.Millisecond + +// MaxConnIDLen is the maximum length of the connection ID +const MaxConnIDLen = 20 + +// InvalidPacketLimitAES is the maximum number of packets that we can fail to decrypt when using +// AEAD_AES_128_GCM or AEAD_AES_265_GCM. +const InvalidPacketLimitAES = 1 << 52 + +// InvalidPacketLimitChaCha is the maximum number of packets that we can fail to decrypt when using AEAD_CHACHA20_POLY1305. +const InvalidPacketLimitChaCha = 1 << 36 diff --git a/third_party/quic-go/internal/protocol/protocol_test.go b/third_party/quic-go/internal/protocol/protocol_test.go new file mode 100644 index 0000000..edf1895 --- /dev/null +++ b/third_party/quic-go/internal/protocol/protocol_test.go @@ -0,0 +1,39 @@ +package protocol + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestLongHeaderPacketTypeStringer(t *testing.T) { + require.Equal(t, "Initial", PacketTypeInitial.String()) + require.Equal(t, "Retry", PacketTypeRetry.String()) + require.Equal(t, "Handshake", PacketTypeHandshake.String()) + require.Equal(t, "0-RTT Protected", PacketType0RTT.String()) + require.Equal(t, "unknown packet type: 10", PacketType(10).String()) +} + +func TestECNFromIPHeader(t *testing.T) { + require.Equal(t, ECNNon, ParseECNHeaderBits(0)) + require.Equal(t, ECT0, ParseECNHeaderBits(0b00000010)) + require.Equal(t, ECT1, ParseECNHeaderBits(0b00000001)) + require.Equal(t, ECNCE, ParseECNHeaderBits(0b00000011)) + require.Panics(t, func() { ParseECNHeaderBits(0b1010101) }) +} + +func TestECNConversionToIPHeaderBits(t *testing.T) { + for _, v := range [...]ECN{ECNNon, ECT0, ECT1, ECNCE} { + require.Equal(t, v, ParseECNHeaderBits(v.ToHeaderBits())) + } + require.Panics(t, func() { ECN(42).ToHeaderBits() }) +} + +func TestECNStringer(t *testing.T) { + require.Equal(t, "ECN unsupported", ECNUnsupported.String()) + require.Equal(t, "Not-ECT", ECNNon.String()) + require.Equal(t, "ECT(0)", ECT0.String()) + require.Equal(t, "ECT(1)", ECT1.String()) + require.Equal(t, "CE", ECNCE.String()) + require.Equal(t, "invalid ECN value: 42", ECN(42).String()) +} diff --git a/third_party/quic-go/internal/protocol/stream.go b/third_party/quic-go/internal/protocol/stream.go new file mode 100644 index 0000000..5e9cb9b --- /dev/null +++ b/third_party/quic-go/internal/protocol/stream.go @@ -0,0 +1,102 @@ +package protocol + +import "github.com/apernet/quic-go/quicvarint" + +// StreamType encodes if this is a unidirectional or bidirectional stream +type StreamType uint8 + +const ( + // StreamTypeUni is a unidirectional stream + StreamTypeUni StreamType = iota + // StreamTypeBidi is a bidirectional stream + StreamTypeBidi +) + +// InvalidPacketNumber is a stream ID that is invalid. +// The first valid stream ID in QUIC is 0. +const InvalidStreamID StreamID = -1 + +// StreamNum is the stream number +type StreamNum int64 + +const ( + // InvalidStreamNum is an invalid stream number. + InvalidStreamNum = -1 + // MaxStreamCount is the maximum stream count value that can be sent in MAX_STREAMS frames + // and as the stream count in the transport parameters + MaxStreamCount StreamNum = 1 << 60 + // MaxStreamID is the maximum stream ID + MaxStreamID StreamID = quicvarint.Max +) + +const ( + // FirstOutgoingBidiStreamClient is the first bidirectional stream opened by the client + FirstOutgoingBidiStreamClient StreamID = 0 + // FirstOutgoingUniStreamClient is the first unidirectional stream opened by the client + FirstOutgoingUniStreamClient StreamID = 2 + // FirstOutgoingBidiStreamServer is the first bidirectional stream opened by the server + FirstOutgoingBidiStreamServer StreamID = 1 + // FirstOutgoingUniStreamServer is the first unidirectional stream opened by the server + FirstOutgoingUniStreamServer StreamID = 3 +) + +const ( + // FirstIncomingBidiStreamServer is the first bidirectional stream accepted by the server + FirstIncomingBidiStreamServer = FirstOutgoingBidiStreamClient + // FirstIncomingUniStreamServer is the first unidirectional stream accepted by the server + FirstIncomingUniStreamServer = FirstOutgoingUniStreamClient + // FirstIncomingBidiStreamClient is the first bidirectional stream accepted by the client + FirstIncomingBidiStreamClient = FirstOutgoingBidiStreamServer + // FirstIncomingUniStreamClient is the first unidirectional stream accepted by the client + FirstIncomingUniStreamClient = FirstOutgoingUniStreamServer +) + +// StreamID calculates the stream ID. +func (s StreamNum) StreamID(stype StreamType, pers Perspective) StreamID { + if s == 0 { + return InvalidStreamID + } + var first StreamID + switch stype { + case StreamTypeBidi: + switch pers { + case PerspectiveClient: + first = 0 + case PerspectiveServer: + first = 1 + } + case StreamTypeUni: + switch pers { + case PerspectiveClient: + first = 2 + case PerspectiveServer: + first = 3 + } + } + return first + 4*StreamID(s-1) +} + +// A StreamID in QUIC +type StreamID int64 + +// StreamInitiator says if the stream was initiated by the client or by the server. +func StreamInitiator(id StreamID) Perspective { + if id%2 == 0 { + return PerspectiveClient + } + return PerspectiveServer +} + +// StreamTypeOf says if this is a unidirectional or bidirectional stream. +func StreamTypeOf(id StreamID) StreamType { + if id%4 >= 2 { + return StreamTypeUni + } + return StreamTypeBidi +} + +// StreamNum returns how many streams in total are below this +// Example: for stream 9 it returns 3 (i.e. streams 1, 5 and 9) +func (s StreamID) StreamNum() StreamNum { + return StreamNum(s/4) + 1 +} diff --git a/third_party/quic-go/internal/protocol/stream_test.go b/third_party/quic-go/internal/protocol/stream_test.go new file mode 100644 index 0000000..916a0b6 --- /dev/null +++ b/third_party/quic-go/internal/protocol/stream_test.go @@ -0,0 +1,66 @@ +package protocol + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestInvalidStreamIDSmallerThanAllValidStreamIDs(t *testing.T) { + require.Less(t, InvalidStreamID, StreamID(0)) +} + +func TestStreamInitiator(t *testing.T) { + require.Equal(t, PerspectiveClient, StreamInitiator(4)) + require.Equal(t, PerspectiveServer, StreamInitiator(5)) + require.Equal(t, PerspectiveClient, StreamInitiator(6)) + require.Equal(t, PerspectiveServer, StreamInitiator(7)) +} + +func TestStreamTypeOf(t *testing.T) { + require.Equal(t, StreamTypeBidi, StreamTypeOf(4)) + require.Equal(t, StreamTypeBidi, StreamTypeOf(5)) + require.Equal(t, StreamTypeUni, StreamTypeOf(6)) + require.Equal(t, StreamTypeUni, StreamTypeOf(7)) +} + +func TestStreamIDStreamNum(t *testing.T) { + require.Equal(t, StreamNum(1), StreamID(0).StreamNum()) + require.Equal(t, StreamNum(1), StreamID(1).StreamNum()) + require.Equal(t, StreamNum(1), StreamID(2).StreamNum()) + require.Equal(t, StreamNum(1), StreamID(3).StreamNum()) + require.Equal(t, StreamNum(3), StreamID(8).StreamNum()) + require.Equal(t, StreamNum(3), StreamID(9).StreamNum()) + require.Equal(t, StreamNum(3), StreamID(10).StreamNum()) + require.Equal(t, StreamNum(3), StreamID(11).StreamNum()) +} + +func TestStreamIDNumToStreamID(t *testing.T) { + // 1st stream + require.Equal(t, StreamID(0), StreamNum(1).StreamID(StreamTypeBidi, PerspectiveClient)) + require.Equal(t, StreamID(1), StreamNum(1).StreamID(StreamTypeBidi, PerspectiveServer)) + require.Equal(t, StreamID(2), StreamNum(1).StreamID(StreamTypeUni, PerspectiveClient)) + require.Equal(t, StreamID(3), StreamNum(1).StreamID(StreamTypeUni, PerspectiveServer)) + + // 100th stream + require.Equal(t, StreamID(396), StreamNum(100).StreamID(StreamTypeBidi, PerspectiveClient)) + require.Equal(t, StreamID(397), StreamNum(100).StreamID(StreamTypeBidi, PerspectiveServer)) + require.Equal(t, StreamID(398), StreamNum(100).StreamID(StreamTypeUni, PerspectiveClient)) + require.Equal(t, StreamID(399), StreamNum(100).StreamID(StreamTypeUni, PerspectiveServer)) + + // 0 is not a valid stream number + require.Equal(t, InvalidStreamID, StreamNum(0).StreamID(StreamTypeBidi, PerspectiveClient)) + require.Equal(t, InvalidStreamID, StreamNum(0).StreamID(StreamTypeBidi, PerspectiveServer)) + require.Equal(t, InvalidStreamID, StreamNum(0).StreamID(StreamTypeUni, PerspectiveClient)) + require.Equal(t, InvalidStreamID, StreamNum(0).StreamID(StreamTypeUni, PerspectiveServer)) +} + +func TestMaxStreamCountValue(t *testing.T) { + const maxStreamID = StreamID(1<<62 - 1) + for _, dir := range []StreamType{StreamTypeUni, StreamTypeBidi} { + for _, pers := range []Perspective{PerspectiveClient, PerspectiveServer} { + require.LessOrEqual(t, MaxStreamCount.StreamID(dir, pers), maxStreamID) + require.Greater(t, (MaxStreamCount+1).StreamID(dir, pers), maxStreamID) + } + } +} diff --git a/third_party/quic-go/internal/protocol/version.go b/third_party/quic-go/internal/protocol/version.go new file mode 100644 index 0000000..8abca50 --- /dev/null +++ b/third_party/quic-go/internal/protocol/version.go @@ -0,0 +1,115 @@ +package protocol + +import ( + "crypto/rand" + "encoding/binary" + "fmt" + "math" + mrand "math/rand/v2" + "slices" + "sync" +) + +// Version is a version number as int +type Version uint32 + +// gQUIC version range as defined in the wiki: https://github.com/quicwg/base-drafts/wiki/QUIC-Versions +const ( + gquicVersion0 = 0x51303030 + maxGquicVersion = 0x51303439 +) + +// The version numbers, making grepping easier +const ( + VersionUnknown Version = math.MaxUint32 + versionDraft29 Version = 0xff00001d // draft-29 used to be a widely deployed version + Version1 Version = 0x1 + Version2 Version = 0x6b3343cf +) + +// SupportedVersions lists the versions that the server supports +// must be in sorted descending order +var SupportedVersions = []Version{Version1, Version2} + +// IsValidVersion says if the version is known to quic-go +func IsValidVersion(v Version) bool { + return v == Version1 || IsSupportedVersion(SupportedVersions, v) +} + +func (vn Version) String() string { + switch vn { + case VersionUnknown: + return "unknown" + case versionDraft29: + return "draft-29" + case Version1: + return "v1" + case Version2: + return "v2" + default: + if vn.isGQUIC() { + return fmt.Sprintf("gQUIC %d", vn.toGQUICVersion()) + } + return fmt.Sprintf("%#x", uint32(vn)) + } +} + +func (vn Version) isGQUIC() bool { + return vn > gquicVersion0 && vn <= maxGquicVersion +} + +func (vn Version) toGQUICVersion() int { + return int(10*(vn-gquicVersion0)/0x100) + int(vn%0x10) +} + +// IsSupportedVersion returns true if the server supports this version +func IsSupportedVersion(supported []Version, v Version) bool { + return slices.Contains(supported, v) +} + +// ChooseSupportedVersion finds the best version in the overlap of ours and theirs +// ours is a slice of versions that we support, sorted by our preference (descending) +// theirs is a slice of versions offered by the peer. The order does not matter. +// The bool returned indicates if a matching version was found. +func ChooseSupportedVersion(ours, theirs []Version) (Version, bool) { + for _, ourVer := range ours { + if slices.Contains(theirs, ourVer) { + return ourVer, true + } + } + return 0, false +} + +var ( + versionNegotiationMx sync.Mutex + versionNegotiationRand mrand.Rand +) + +func init() { + var seed [16]byte + rand.Read(seed[:]) + versionNegotiationRand = *mrand.New(mrand.NewPCG( + binary.BigEndian.Uint64(seed[:8]), + binary.BigEndian.Uint64(seed[8:]), + )) +} + +// generateReservedVersion generates a reserved version (v & 0x0f0f0f0f == 0x0a0a0a0a) +func generateReservedVersion() Version { + var b [4]byte + binary.BigEndian.PutUint32(b[:], versionNegotiationRand.Uint32()) + return Version((binary.BigEndian.Uint32(b[:]) | 0x0a0a0a0a) & 0xfafafafa) +} + +// GetGreasedVersions adds one reserved version number to a slice of version numbers, at a random position. +// It doesn't modify the supported slice. +func GetGreasedVersions(supported []Version) []Version { + versionNegotiationMx.Lock() + defer versionNegotiationMx.Unlock() + randPos := versionNegotiationRand.IntN(len(supported) + 1) + greased := make([]Version, len(supported)+1) + copy(greased, supported[:randPos]) + greased[randPos] = generateReservedVersion() + copy(greased[randPos+1:], supported[randPos:]) + return greased +} diff --git a/third_party/quic-go/internal/protocol/version_test.go b/third_party/quic-go/internal/protocol/version_test.go new file mode 100644 index 0000000..f82137f --- /dev/null +++ b/third_party/quic-go/internal/protocol/version_test.go @@ -0,0 +1,151 @@ +package protocol + +import ( + "slices" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestValidVersion(t *testing.T) { + require.False(t, IsValidVersion(VersionUnknown)) + require.False(t, IsValidVersion(versionDraft29)) + require.True(t, IsValidVersion(Version1)) + require.True(t, IsValidVersion(Version2)) + require.False(t, IsValidVersion(1234)) +} + +func TestVersionStringRepresentation(t *testing.T) { + require.Equal(t, "unknown", VersionUnknown.String()) + require.Equal(t, "draft-29", versionDraft29.String()) + require.Equal(t, "v1", Version1.String()) + require.Equal(t, "v2", Version2.String()) + // check with unsupported version numbers from the wiki + require.Equal(t, "gQUIC 9", Version(0x51303039).String()) + require.Equal(t, "gQUIC 13", Version(0x51303133).String()) + require.Equal(t, "gQUIC 25", Version(0x51303235).String()) + require.Equal(t, "gQUIC 48", Version(0x51303438).String()) + require.Equal(t, "0x1234567", Version(0x01234567).String()) +} + +func TestRecognizesSupportedVersions(t *testing.T) { + require.False(t, IsSupportedVersion(SupportedVersions, 0)) + require.False(t, IsSupportedVersion(SupportedVersions, maxGquicVersion)) + require.True(t, IsSupportedVersion(SupportedVersions, SupportedVersions[0])) + require.True(t, IsSupportedVersion(SupportedVersions, SupportedVersions[len(SupportedVersions)-1])) +} + +func TestVersionSelection(t *testing.T) { + tests := []struct { + name string + supportedVersions []Version + otherVersions []Version + expectedVersion Version + expectedOK bool + }{ + { + name: "finds matching version", + supportedVersions: []Version{1, 2, 3}, + otherVersions: []Version{6, 5, 4, 3}, + expectedVersion: 3, + expectedOK: true, + }, + { + name: "picks preferred version", + supportedVersions: []Version{2, 1, 3}, + otherVersions: []Version{3, 6, 1, 8, 2, 10}, + expectedVersion: 2, + expectedOK: true, + }, + { + name: "no matching version", + supportedVersions: []Version{1}, + otherVersions: []Version{2}, + expectedOK: false, + }, + { + name: "empty supported versions", + supportedVersions: []Version{}, + otherVersions: []Version{1, 2}, + expectedOK: false, + }, + { + name: "empty other versions", + supportedVersions: []Version{102, 101}, + otherVersions: []Version{}, + expectedOK: false, + }, + { + name: "both empty", + supportedVersions: []Version{}, + otherVersions: []Version{}, + expectedOK: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ver, ok := ChooseSupportedVersion(tt.supportedVersions, tt.otherVersions) + require.Equal(t, tt.expectedOK, ok) + if tt.expectedOK { + require.Equal(t, tt.expectedVersion, ver) + } + }) + } +} + +func isReservedVersion(v Version) bool { return v&0x0f0f0f0f == 0x0a0a0a0a } + +func TestVersionGreasing(t *testing.T) { + // adding to an empty slice + greased := GetGreasedVersions([]Version{}) + require.Len(t, greased, 1) + require.True(t, isReservedVersion(greased[0])) + + // make sure that the greased versions are distinct, + // allowing for a small number of duplicates + var versions []Version + for range 25 { + versions = GetGreasedVersions(versions) + } + slices.Sort(versions) + var numDuplicates int + for i, v := range versions { + require.True(t, isReservedVersion(v)) + if i > 0 && versions[i-1] == v { + numDuplicates++ + } + } + require.LessOrEqual(t, numDuplicates, 3) + + // adding it somewhere in a slice of supported versions + supported := []Version{10, 18, 29} + for _, v := range supported { + require.False(t, isReservedVersion(v)) + } + + var greasedVersionFirst, greasedVersionLast, greasedVersionMiddle int + for range 100 { + greased := GetGreasedVersions(supported) + require.Len(t, greased, 4) + + var j int + for i, v := range greased { + if isReservedVersion(v) { + if i == 0 { + greasedVersionFirst++ + } + if i == len(greased)-1 { + greasedVersionLast++ + } + greasedVersionMiddle++ + continue + } + require.Equal(t, supported[j], v) + j++ + } + } + require.NotZero(t, greasedVersionFirst) + require.NotZero(t, greasedVersionLast) + require.NotZero(t, greasedVersionMiddle) +} diff --git a/third_party/quic-go/internal/qerr/error_codes.go b/third_party/quic-go/internal/qerr/error_codes.go new file mode 100644 index 0000000..0036130 --- /dev/null +++ b/third_party/quic-go/internal/qerr/error_codes.go @@ -0,0 +1,87 @@ +package qerr + +import ( + "crypto/tls" + "fmt" +) + +// TransportErrorCode is a QUIC transport error. +type TransportErrorCode uint64 + +// The error codes defined by QUIC +const ( + NoError TransportErrorCode = 0x0 + InternalError TransportErrorCode = 0x1 + ConnectionRefused TransportErrorCode = 0x2 + FlowControlError TransportErrorCode = 0x3 + StreamLimitError TransportErrorCode = 0x4 + StreamStateError TransportErrorCode = 0x5 + FinalSizeError TransportErrorCode = 0x6 + FrameEncodingError TransportErrorCode = 0x7 + TransportParameterError TransportErrorCode = 0x8 + ConnectionIDLimitError TransportErrorCode = 0x9 + ProtocolViolation TransportErrorCode = 0xa + InvalidToken TransportErrorCode = 0xb + ApplicationErrorErrorCode TransportErrorCode = 0xc + CryptoBufferExceeded TransportErrorCode = 0xd + KeyUpdateError TransportErrorCode = 0xe + AEADLimitReached TransportErrorCode = 0xf + NoViablePathError TransportErrorCode = 0x10 +) + +func (e TransportErrorCode) IsCryptoError() bool { + return e >= 0x100 && e < 0x200 +} + +// Message is a description of the error. +// It only returns a non-empty string for crypto errors. +func (e TransportErrorCode) Message() string { + if !e.IsCryptoError() { + return "" + } + return tls.AlertError(e - 0x100).Error() +} + +func (e TransportErrorCode) String() string { + switch e { + case NoError: + return "NO_ERROR" + case InternalError: + return "INTERNAL_ERROR" + case ConnectionRefused: + return "CONNECTION_REFUSED" + case FlowControlError: + return "FLOW_CONTROL_ERROR" + case StreamLimitError: + return "STREAM_LIMIT_ERROR" + case StreamStateError: + return "STREAM_STATE_ERROR" + case FinalSizeError: + return "FINAL_SIZE_ERROR" + case FrameEncodingError: + return "FRAME_ENCODING_ERROR" + case TransportParameterError: + return "TRANSPORT_PARAMETER_ERROR" + case ConnectionIDLimitError: + return "CONNECTION_ID_LIMIT_ERROR" + case ProtocolViolation: + return "PROTOCOL_VIOLATION" + case InvalidToken: + return "INVALID_TOKEN" + case ApplicationErrorErrorCode: + return "APPLICATION_ERROR" + case CryptoBufferExceeded: + return "CRYPTO_BUFFER_EXCEEDED" + case KeyUpdateError: + return "KEY_UPDATE_ERROR" + case AEADLimitReached: + return "AEAD_LIMIT_REACHED" + case NoViablePathError: + return "NO_VIABLE_PATH" + default: + if e.IsCryptoError() { + return fmt.Sprintf("CRYPTO_ERROR %#x", uint16(e)) + } + return fmt.Sprintf("unknown error code: %#x", uint16(e)) + } +} diff --git a/third_party/quic-go/internal/qerr/errorcodes_test.go b/third_party/quic-go/internal/qerr/errorcodes_test.go new file mode 100644 index 0000000..2e07fd4 --- /dev/null +++ b/third_party/quic-go/internal/qerr/errorcodes_test.go @@ -0,0 +1,47 @@ +package qerr + +import ( + "go/ast" + "go/parser" + "go/token" + "path" + "runtime" + "strconv" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestTransportErrorCodeStringer(t *testing.T) { + _, thisfile, _, ok := runtime.Caller(0) + require.True(t, ok, "Failed to get current frame") + + filename := path.Join(path.Dir(thisfile), "error_codes.go") + fileAst, err := parser.ParseFile(token.NewFileSet(), filename, nil, 0) + require.NoError(t, err) + + constSpecs := fileAst.Decls[2].(*ast.GenDecl).Specs + require.Greater(t, len(constSpecs), 4, "Expected more than 4 constants") + + for _, c := range constSpecs { + valString := c.(*ast.ValueSpec).Values[0].(*ast.BasicLit).Value + val, err := strconv.ParseInt(valString, 0, 64) + require.NoError(t, err) + require.NotEqual(t, "unknown error code", TransportErrorCode(val).String()) + } + + // test that there's a string representation for unknown error codes + require.Equal(t, "unknown error code: 0x1337", TransportErrorCode(0x1337).String()) +} + +func TestIsCryptoError(t *testing.T) { + for i := range 0x100 { + require.False(t, TransportErrorCode(i).IsCryptoError()) + } + for i := 0x100; i < 0x200; i++ { + require.True(t, TransportErrorCode(i).IsCryptoError()) + } + for i := 0x200; i < 0x300; i++ { + require.False(t, TransportErrorCode(i).IsCryptoError()) + } +} diff --git a/third_party/quic-go/internal/qerr/errors.go b/third_party/quic-go/internal/qerr/errors.go new file mode 100644 index 0000000..eb3097e --- /dev/null +++ b/third_party/quic-go/internal/qerr/errors.go @@ -0,0 +1,134 @@ +package qerr + +import ( + "fmt" + "net" + + "github.com/apernet/quic-go/internal/protocol" +) + +var ( + ErrHandshakeTimeout = &HandshakeTimeoutError{} + ErrIdleTimeout = &IdleTimeoutError{} +) + +type TransportError struct { + Remote bool + FrameType uint64 + ErrorCode TransportErrorCode + ErrorMessage string + error error // only set for local errors, sometimes +} + +var _ error = &TransportError{} + +// NewLocalCryptoError create a new TransportError instance for a crypto error +func NewLocalCryptoError(tlsAlert uint8, err error) *TransportError { + return &TransportError{ + ErrorCode: 0x100 + TransportErrorCode(tlsAlert), + error: err, + } +} + +func (e *TransportError) Error() string { + str := fmt.Sprintf("%s (%s)", e.ErrorCode.String(), getRole(e.Remote)) + if e.FrameType != 0 { + str += fmt.Sprintf(" (frame type: %#x)", e.FrameType) + } + msg := e.ErrorMessage + if len(msg) == 0 && e.error != nil { + msg = e.error.Error() + } + if len(msg) == 0 { + msg = e.ErrorCode.Message() + } + if len(msg) == 0 { + return str + } + return str + ": " + msg +} + +func (e *TransportError) Unwrap() []error { return []error{net.ErrClosed, e.error} } + +func (e *TransportError) Is(target error) bool { + t, ok := target.(*TransportError) + return ok && e.ErrorCode == t.ErrorCode && e.FrameType == t.FrameType && e.Remote == t.Remote +} + +// An ApplicationErrorCode is an application-defined error code. +type ApplicationErrorCode uint64 + +// A StreamErrorCode is an error code used to cancel streams. +type StreamErrorCode uint64 + +type ApplicationError struct { + Remote bool + ErrorCode ApplicationErrorCode + ErrorMessage string +} + +var _ error = &ApplicationError{} + +func (e *ApplicationError) Error() string { + if len(e.ErrorMessage) == 0 { + return fmt.Sprintf("Application error %#x (%s)", e.ErrorCode, getRole(e.Remote)) + } + return fmt.Sprintf("Application error %#x (%s): %s", e.ErrorCode, getRole(e.Remote), e.ErrorMessage) +} + +func (e *ApplicationError) Unwrap() error { return net.ErrClosed } + +func (e *ApplicationError) Is(target error) bool { + t, ok := target.(*ApplicationError) + return ok && e.ErrorCode == t.ErrorCode && e.Remote == t.Remote +} + +type IdleTimeoutError struct{} + +var _ error = &IdleTimeoutError{} + +func (e *IdleTimeoutError) Timeout() bool { return true } +func (e *IdleTimeoutError) Temporary() bool { return false } +func (e *IdleTimeoutError) Error() string { return "timeout: no recent network activity" } +func (e *IdleTimeoutError) Unwrap() error { return net.ErrClosed } + +type HandshakeTimeoutError struct{} + +var _ error = &HandshakeTimeoutError{} + +func (e *HandshakeTimeoutError) Timeout() bool { return true } +func (e *HandshakeTimeoutError) Temporary() bool { return false } +func (e *HandshakeTimeoutError) Error() string { return "timeout: handshake did not complete in time" } +func (e *HandshakeTimeoutError) Unwrap() error { return net.ErrClosed } + +// A VersionNegotiationError occurs when the client and the server can't agree on a QUIC version. +type VersionNegotiationError struct { + Ours []protocol.Version + Theirs []protocol.Version +} + +func (e *VersionNegotiationError) Error() string { + return fmt.Sprintf("no compatible QUIC version found (we support %s, server offered %s)", e.Ours, e.Theirs) +} + +func (e *VersionNegotiationError) Unwrap() error { return net.ErrClosed } + +// A StatelessResetError occurs when we receive a stateless reset. +type StatelessResetError struct{} + +var _ net.Error = &StatelessResetError{} + +func (e *StatelessResetError) Error() string { + return "received a stateless reset" +} + +func (e *StatelessResetError) Unwrap() error { return net.ErrClosed } +func (e *StatelessResetError) Timeout() bool { return false } +func (e *StatelessResetError) Temporary() bool { return true } + +func getRole(remote bool) string { + if remote { + return "remote" + } + return "local" +} diff --git a/third_party/quic-go/internal/qerr/errors_test.go b/third_party/quic-go/internal/qerr/errors_test.go new file mode 100644 index 0000000..fa08c61 --- /dev/null +++ b/third_party/quic-go/internal/qerr/errors_test.go @@ -0,0 +1,176 @@ +package qerr + +import ( + "errors" + "fmt" + "net" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestTransportError(t *testing.T) { + require.True(t, errors.Is(&TransportError{}, net.ErrClosed)) + + require.True(t, errors.Is( + &TransportError{Remote: true, ErrorCode: 1, FrameType: 2}, + &TransportError{Remote: true, ErrorCode: 1, FrameType: 2}, + )) + require.True(t, errors.Is(&TransportError{ErrorCode: 0x101}, &TransportError{ErrorCode: 0x101})) + require.False(t, errors.Is(&TransportError{}, &TransportError{ErrorCode: 0x101})) + require.False(t, errors.Is(&TransportError{}, &TransportError{FrameType: 0x1337})) + require.False(t, errors.Is(&TransportError{Remote: false}, &TransportError{Remote: true})) +} + +func TestTransportErrorStringer(t *testing.T) { + t.Run("with error message", func(t *testing.T) { + err := &TransportError{ + ErrorCode: FlowControlError, + ErrorMessage: "foobar", + } + require.Equal(t, "FLOW_CONTROL_ERROR (local): foobar", err.Error()) + }) + + t.Run("without error message", func(t *testing.T) { + err := &TransportError{ErrorCode: FlowControlError} + require.Equal(t, "FLOW_CONTROL_ERROR (local)", err.Error()) + }) + + t.Run("with frame type", func(t *testing.T) { + err := &TransportError{ + Remote: true, + ErrorCode: FlowControlError, + FrameType: 0x1337, + } + require.Equal(t, "FLOW_CONTROL_ERROR (remote) (frame type: 0x1337)", err.Error()) + }) + + t.Run("with frame type and error message", func(t *testing.T) { + err := &TransportError{ + ErrorCode: FlowControlError, + FrameType: 0x1337, + ErrorMessage: "foobar", + } + require.Equal(t, "FLOW_CONTROL_ERROR (local) (frame type: 0x1337): foobar", err.Error()) + }) +} + +type myError int + +var _ error = myError(0) + +func (e myError) Error() string { return fmt.Sprintf("my error %d", e) } + +func TestCryptoError(t *testing.T) { + var myErr myError + err := NewLocalCryptoError(0x42, myError(1337)) + require.True(t, errors.As(err, &myErr)) + require.Equal(t, myError(1337), myErr) + + err = NewLocalCryptoError(0x42, assert.AnError) + require.True(t, errors.Is(err, assert.AnError)) + require.True(t, errors.Is( + NewLocalCryptoError(0x42, assert.AnError), + NewLocalCryptoError(0x42, assert.AnError), + )) + require.False(t, errors.Is( + NewLocalCryptoError(0x42, assert.AnError), + NewLocalCryptoError(0x43, assert.AnError), + )) +} + +func TestCryptoErrorStringer(t *testing.T) { + t.Run("with error message", func(t *testing.T) { + myErr := myError(1337) + err := NewLocalCryptoError(0x42, myErr) + require.Equal(t, "CRYPTO_ERROR 0x142 (local): my error 1337", err.Error()) + }) + + t.Run("without error message", func(t *testing.T) { + err := NewLocalCryptoError(0x2a, nil) + require.Equal(t, "CRYPTO_ERROR 0x12a (local): tls: bad certificate", err.Error()) + }) +} + +func TestApplicationError(t *testing.T) { + require.True(t, errors.Is(&ApplicationError{}, net.ErrClosed)) + + require.True(t, errors.Is( + &ApplicationError{ErrorCode: 1, Remote: true}, + &ApplicationError{ErrorCode: 1, Remote: true}, + )) + require.True(t, errors.Is(&ApplicationError{ErrorCode: 0x101}, &ApplicationError{ErrorCode: 0x101})) + require.False(t, errors.Is(&ApplicationError{}, &ApplicationError{ErrorCode: 0x101})) + require.False(t, errors.Is(&ApplicationError{Remote: false}, &ApplicationError{Remote: true})) +} + +func TestApplicationErrorStringer(t *testing.T) { + t.Run("with error message", func(t *testing.T) { + err := &ApplicationError{ + ErrorCode: 0x42, + ErrorMessage: "foobar", + } + require.Equal(t, "Application error 0x42 (local): foobar", err.Error()) + }) + + t.Run("without error message", func(t *testing.T) { + err := &ApplicationError{ + ErrorCode: 0x42, + Remote: true, + } + require.Equal(t, "Application error 0x42 (remote)", err.Error()) + }) +} + +func TestHandshakeTimeoutError(t *testing.T) { + require.True(t, errors.Is(&HandshakeTimeoutError{}, &HandshakeTimeoutError{})) + require.False(t, errors.Is(&HandshakeTimeoutError{}, &IdleTimeoutError{})) + + //nolint:staticcheck // SA1021: we need to assign to an interface here + var err error + err = &HandshakeTimeoutError{} + nerr, ok := err.(net.Error) + require.True(t, ok) + require.True(t, nerr.Timeout()) + require.Equal(t, "timeout: handshake did not complete in time", err.Error()) + require.True(t, errors.Is(&HandshakeTimeoutError{}, net.ErrClosed)) +} + +func TestIdleTimeoutError(t *testing.T) { + require.True(t, errors.Is(&IdleTimeoutError{}, &IdleTimeoutError{})) + require.False(t, errors.Is(&IdleTimeoutError{}, &HandshakeTimeoutError{})) + + //nolint:staticcheck // SA1021: we need to assign to an interface here + var err error + err = &IdleTimeoutError{} + nerr, ok := err.(net.Error) + require.True(t, ok) + require.True(t, nerr.Timeout()) + require.Equal(t, "timeout: no recent network activity", err.Error()) + require.True(t, errors.Is(&IdleTimeoutError{}, net.ErrClosed)) +} + +func TestVersionNegotiationErrorString(t *testing.T) { + err := &VersionNegotiationError{ + Ours: []protocol.Version{2, 3}, + Theirs: []protocol.Version{4, 5, 6}, + } + require.Equal(t, "no compatible QUIC version found (we support [0x2 0x3], server offered [0x4 0x5 0x6])", err.Error()) + require.True(t, errors.Is(&VersionNegotiationError{}, net.ErrClosed)) +} + +func TestStatelessResetError(t *testing.T) { + require.Equal(t, "received a stateless reset", (&StatelessResetError{}).Error()) + require.True(t, errors.Is(&StatelessResetError{}, &StatelessResetError{})) + + //nolint:staticcheck // SA1021: we need to assign to an interface here + var err error + err = &StatelessResetError{} + nerr, ok := err.(net.Error) + require.True(t, ok) + require.False(t, nerr.Timeout()) + require.True(t, errors.Is(&StatelessResetError{}, net.ErrClosed)) +} diff --git a/third_party/quic-go/internal/qtls/cipher_suite.go b/third_party/quic-go/internal/qtls/cipher_suite.go new file mode 100644 index 0000000..32a921c --- /dev/null +++ b/third_party/quic-go/internal/qtls/cipher_suite.go @@ -0,0 +1,52 @@ +package qtls + +import ( + "crypto/tls" + "fmt" + "unsafe" +) + +//go:linkname cipherSuitesTLS13 crypto/tls.cipherSuitesTLS13 +var cipherSuitesTLS13 []unsafe.Pointer + +//go:linkname defaultCipherSuitesTLS13 crypto/tls.defaultCipherSuitesTLS13 +var defaultCipherSuitesTLS13 []uint16 + +//go:linkname defaultCipherSuitesTLS13NoAES crypto/tls.defaultCipherSuitesTLS13NoAES +var defaultCipherSuitesTLS13NoAES []uint16 + +var cipherSuitesModified bool + +// SetCipherSuite modifies the cipherSuiteTLS13 slice of cipher suites inside qtls +// such that it only contains the cipher suite with the chosen id. +// The reset function returned resets them back to the original value. +func SetCipherSuite(id uint16) (reset func()) { + if cipherSuitesModified { + panic("cipher suites modified multiple times without resetting") + } + cipherSuitesModified = true + + origCipherSuitesTLS13 := append([]unsafe.Pointer{}, cipherSuitesTLS13...) + origDefaultCipherSuitesTLS13 := append([]uint16{}, defaultCipherSuitesTLS13...) + origDefaultCipherSuitesTLS13NoAES := append([]uint16{}, defaultCipherSuitesTLS13NoAES...) + // The order is given by the order of the slice elements in cipherSuitesTLS13 in qtls. + switch id { + case tls.TLS_AES_128_GCM_SHA256: + cipherSuitesTLS13 = cipherSuitesTLS13[:1] + case tls.TLS_CHACHA20_POLY1305_SHA256: + cipherSuitesTLS13 = cipherSuitesTLS13[1:2] + case tls.TLS_AES_256_GCM_SHA384: + cipherSuitesTLS13 = cipherSuitesTLS13[2:] + default: + panic(fmt.Sprintf("unexpected cipher suite: %d", id)) + } + defaultCipherSuitesTLS13 = []uint16{id} + defaultCipherSuitesTLS13NoAES = []uint16{id} + + return func() { + cipherSuitesTLS13 = origCipherSuitesTLS13 + defaultCipherSuitesTLS13 = origDefaultCipherSuitesTLS13 + defaultCipherSuitesTLS13NoAES = origDefaultCipherSuitesTLS13NoAES + cipherSuitesModified = false + } +} diff --git a/third_party/quic-go/internal/qtls/cipher_suite_test.go b/third_party/quic-go/internal/qtls/cipher_suite_test.go new file mode 100644 index 0000000..9b54bc8 --- /dev/null +++ b/third_party/quic-go/internal/qtls/cipher_suite_test.go @@ -0,0 +1,54 @@ +package qtls + +import ( + "crypto/fips140" + "crypto/tls" + "fmt" + "net" + "testing" + + "github.com/apernet/quic-go/internal/testdata" + + "github.com/stretchr/testify/require" +) + +func TestCipherSuiteSelection(t *testing.T) { + t.Run("TLS_AES_128_GCM_SHA256", func(t *testing.T) { testCipherSuiteSelection(t, tls.TLS_AES_128_GCM_SHA256) }) + t.Run("TLS_CHACHA20_POLY1305_SHA256", func(t *testing.T) { testCipherSuiteSelection(t, tls.TLS_CHACHA20_POLY1305_SHA256) }) + t.Run("TLS_AES_256_GCM_SHA384", func(t *testing.T) { testCipherSuiteSelection(t, tls.TLS_AES_256_GCM_SHA384) }) +} + +func testCipherSuiteSelection(t *testing.T, cs uint16) { + if fips140.Enabled() && cs == tls.TLS_CHACHA20_POLY1305_SHA256 { + t.Skip("ChaCha20-Poly1305 is not allowed in FIPS 140-3 mode") + } + + reset := SetCipherSuite(cs) + defer reset() + + ln, err := tls.Listen("tcp4", "localhost:0", testdata.GetTLSConfig()) + require.NoError(t, err) + defer ln.Close() + + done := make(chan struct{}) + go func() { + defer close(done) + conn, err := ln.Accept() + require.NoError(t, err) + _, err = conn.Read(make([]byte, 10)) + require.NoError(t, err) + require.Equal(t, cs, conn.(*tls.Conn).ConnectionState().CipherSuite) + }() + + conn, err := tls.Dial( + "tcp4", + fmt.Sprintf("localhost:%d", ln.Addr().(*net.TCPAddr).Port), + &tls.Config{RootCAs: testdata.GetRootCA()}, + ) + require.NoError(t, err) + _, err = conn.Write([]byte("foobar")) + require.NoError(t, err) + require.Equal(t, cs, conn.ConnectionState().CipherSuite) + require.NoError(t, conn.Close()) + <-done +} diff --git a/third_party/quic-go/internal/testdata/cert.go b/third_party/quic-go/internal/testdata/cert.go new file mode 100644 index 0000000..d7be4b5 --- /dev/null +++ b/third_party/quic-go/internal/testdata/cert.go @@ -0,0 +1,143 @@ +package testdata + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "fmt" + "math/big" + "net" + "os" + "path/filepath" + "sync" + "time" +) + +var ( + certificateOnce sync.Once + certificatePath string + privateKeyPath string + rootCAPEM []byte + certificateErr error +) + +// GetCertificatePaths returns the paths to certificate and key +func GetCertificatePaths() (string, string) { + ensureCertificate() + return certificatePath, privateKeyPath +} + +// GetTLSConfig returns a TLS config for localhost. +func GetTLSConfig() *tls.Config { + cert, err := tls.LoadX509KeyPair(GetCertificatePaths()) + if err != nil { + panic(err) + } + return &tls.Config{ + MinVersion: tls.VersionTLS13, + Certificates: []tls.Certificate{cert}, + } +} + +// AddRootCA adds the root CA certificate to a cert pool +func AddRootCA(certPool *x509.CertPool) { + ensureCertificate() + if ok := certPool.AppendCertsFromPEM(rootCAPEM); !ok { + panic("could not add root certificate to pool") + } +} + +// GetRootCA returns an x509.CertPool containing (only) the CA certificate +func GetRootCA() *x509.CertPool { + pool := x509.NewCertPool() + AddRootCA(pool) + return pool +} + +func ensureCertificate() { + certificateOnce.Do(generateCertificate) + if certificateErr != nil { + panic(certificateErr) + } +} + +func generateCertificate() { + now := time.Now() + caKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + certificateErr = fmt.Errorf("generate test CA key: %w", err) + return + } + caTemplate := &x509.Certificate{ + SerialNumber: randomSerial(), + Subject: pkix.Name{Organization: []string{"quic-go test CA"}}, + NotBefore: now.Add(-time.Minute), + NotAfter: now.Add(24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + BasicConstraintsValid: true, + IsCA: true, + } + caDER, err := x509.CreateCertificate(rand.Reader, caTemplate, caTemplate, &caKey.PublicKey, caKey) + if err != nil { + certificateErr = fmt.Errorf("create test CA certificate: %w", err) + return + } + + leafKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + certificateErr = fmt.Errorf("generate test leaf key: %w", err) + return + } + leafTemplate := &x509.Certificate{ + SerialNumber: randomSerial(), + Subject: pkix.Name{Organization: []string{"quic-go test server"}}, + NotBefore: now.Add(-time.Minute), + NotAfter: now.Add(24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + BasicConstraintsValid: true, + DNSNames: []string{"localhost"}, + IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.IPv6loopback}, + } + leafDER, err := x509.CreateCertificate(rand.Reader, leafTemplate, caTemplate, &leafKey.PublicKey, caKey) + if err != nil { + certificateErr = fmt.Errorf("create test leaf certificate: %w", err) + return + } + keyDER, err := x509.MarshalECPrivateKey(leafKey) + if err != nil { + certificateErr = fmt.Errorf("marshal test leaf key: %w", err) + return + } + rootCAPEM = pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: caDER}) + leafPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: leafDER}) + leafPEM = append(leafPEM, rootCAPEM...) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + tempDir, err := os.MkdirTemp("", "quic-go-test-cert-") + if err != nil { + certificateErr = fmt.Errorf("create test certificate directory: %w", err) + return + } + certificatePath = filepath.Join(tempDir, "cert.pem") + privateKeyPath = filepath.Join(tempDir, "priv.key") + if err := os.WriteFile(certificatePath, leafPEM, 0o600); err != nil { + certificateErr = fmt.Errorf("write test certificate: %w", err) + return + } + if err := os.WriteFile(privateKeyPath, keyPEM, 0o600); err != nil { + certificateErr = fmt.Errorf("write test private key: %w", err) + } +} + +func randomSerial() *big.Int { + limit := new(big.Int).Lsh(big.NewInt(1), 128) + serial, err := rand.Int(rand.Reader, limit) + if err != nil { + panic(fmt.Errorf("generate test certificate serial: %w", err)) + } + return serial +} diff --git a/third_party/quic-go/internal/testdata/cert_test.go b/third_party/quic-go/internal/testdata/cert_test.go new file mode 100644 index 0000000..e3e4a79 --- /dev/null +++ b/third_party/quic-go/internal/testdata/cert_test.go @@ -0,0 +1,31 @@ +package testdata + +import ( + "crypto/tls" + "io" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCertificates(t *testing.T) { + ln, err := tls.Listen("tcp", "localhost:0", GetTLSConfig()) + require.NoError(t, err) + + go func() { + conn, err := ln.Accept() + require.NoError(t, err) + defer conn.Close() + _, err = conn.Write([]byte("foobar")) + require.NoError(t, err) + }() + + conn, err := tls.Dial("tcp", ln.Addr().String(), &tls.Config{ + RootCAs: GetRootCA(), + ServerName: "localhost", + }) + require.NoError(t, err) + data, err := io.ReadAll(conn) + require.NoError(t, err) + require.Equal(t, "foobar", string(data)) +} diff --git a/third_party/quic-go/internal/utils/buffered_write_closer.go b/third_party/quic-go/internal/utils/buffered_write_closer.go new file mode 100644 index 0000000..80ebfec --- /dev/null +++ b/third_party/quic-go/internal/utils/buffered_write_closer.go @@ -0,0 +1,26 @@ +package utils + +import ( + "bufio" + "io" +) + +type bufferedWriteCloser struct { + *bufio.Writer + io.Closer +} + +// NewBufferedWriteCloser creates an io.WriteCloser from a bufio.Writer and an io.Closer +func NewBufferedWriteCloser(writer *bufio.Writer, closer io.Closer) io.WriteCloser { + return &bufferedWriteCloser{ + Writer: writer, + Closer: closer, + } +} + +func (h bufferedWriteCloser) Close() error { + if err := h.Flush(); err != nil { + return err + } + return h.Closer.Close() +} diff --git a/third_party/quic-go/internal/utils/buffered_write_closer_test.go b/third_party/quic-go/internal/utils/buffered_write_closer_test.go new file mode 100644 index 0000000..4abc9d9 --- /dev/null +++ b/third_party/quic-go/internal/utils/buffered_write_closer_test.go @@ -0,0 +1,25 @@ +package utils + +import ( + "bufio" + "bytes" + "testing" + + "github.com/stretchr/testify/require" +) + +type nopCloser struct{} + +func (nopCloser) Close() error { return nil } + +func TestBufferedWriteCloserFlushBeforeClosing(t *testing.T) { + buf := &bytes.Buffer{} + + w := bufio.NewWriter(buf) + wc := NewBufferedWriteCloser(w, &nopCloser{}) + _, err := wc.Write([]byte("foobar")) + require.NoError(t, err) + require.Zero(t, buf.Len()) + require.NoError(t, wc.Close()) + require.Equal(t, "foobar", buf.String()) +} diff --git a/third_party/quic-go/internal/utils/connstats.go b/third_party/quic-go/internal/utils/connstats.go new file mode 100644 index 0000000..19d8831 --- /dev/null +++ b/third_party/quic-go/internal/utils/connstats.go @@ -0,0 +1,14 @@ +package utils + +import "sync/atomic" + +// ConnectionStats stores stats for the connection. See the public +// ConnectionStats struct in connection.go for more information +type ConnectionStats struct { + BytesSent atomic.Uint64 + PacketsSent atomic.Uint64 + BytesReceived atomic.Uint64 + PacketsReceived atomic.Uint64 + BytesLost atomic.Uint64 + PacketsLost atomic.Uint64 +} diff --git a/third_party/quic-go/internal/utils/linkedlist/README.md b/third_party/quic-go/internal/utils/linkedlist/README.md new file mode 100644 index 0000000..42c2172 --- /dev/null +++ b/third_party/quic-go/internal/utils/linkedlist/README.md @@ -0,0 +1,6 @@ +# Usage + +This is the Go standard library implementation of a linked list +(https://golang.org/src/container/list/list.go), with the following modifications: +* it uses Go generics +* it allows passing in a `sync.Pool` (via the `NewWithPool` constructor) to reduce allocations of `Element` structs diff --git a/third_party/quic-go/internal/utils/linkedlist/linkedlist.go b/third_party/quic-go/internal/utils/linkedlist/linkedlist.go new file mode 100644 index 0000000..804a344 --- /dev/null +++ b/third_party/quic-go/internal/utils/linkedlist/linkedlist.go @@ -0,0 +1,264 @@ +// Copyright 2009 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// Package list implements a doubly linked list. +// +// To iterate over a list (where l is a *List[T]): +// +// for e := l.Front(); e != nil; e = e.Next() { +// // do something with e.Value +// } +package list + +import "sync" + +func NewPool[T any]() *sync.Pool { + return &sync.Pool{New: func() any { return &Element[T]{} }} +} + +// Element is an element of a linked list. +type Element[T any] struct { + // Next and previous pointers in the doubly-linked list of elements. + // To simplify the implementation, internally a list l is implemented + // as a ring, such that &l.root is both the next element of the last + // list element (l.Back()) and the previous element of the first list + // element (l.Front()). + next, prev *Element[T] + + // The list to which this element belongs. + list *List[T] + + // The value stored with this element. + Value T +} + +// Next returns the next list element or nil. +func (e *Element[T]) Next() *Element[T] { + if p := e.next; e.list != nil && p != &e.list.root { + return p + } + return nil +} + +// Prev returns the previous list element or nil. +func (e *Element[T]) Prev() *Element[T] { + if p := e.prev; e.list != nil && p != &e.list.root { + return p + } + return nil +} + +func (e *Element[T]) List() *List[T] { + return e.list +} + +// List represents a doubly linked list. +// The zero value for List is an empty list ready to use. +type List[T any] struct { + root Element[T] // sentinel list element, only &root, root.prev, and root.next are used + len int // current list length excluding (this) sentinel element + + pool *sync.Pool +} + +// Init initializes or clears list l. +func (l *List[T]) Init() *List[T] { + l.root.next = &l.root + l.root.prev = &l.root + l.len = 0 + return l +} + +// New returns an initialized list. +func New[T any]() *List[T] { return new(List[T]).Init() } + +// NewWithPool returns an initialized list, using a sync.Pool for list elements. +func NewWithPool[T any](pool *sync.Pool) *List[T] { + l := &List[T]{pool: pool} + return l.Init() +} + +// Len returns the number of elements of list l. +// The complexity is O(1). +func (l *List[T]) Len() int { return l.len } + +// Front returns the first element of list l or nil if the list is empty. +func (l *List[T]) Front() *Element[T] { + if l.len == 0 { + return nil + } + return l.root.next +} + +// Back returns the last element of list l or nil if the list is empty. +func (l *List[T]) Back() *Element[T] { + if l.len == 0 { + return nil + } + return l.root.prev +} + +// lazyInit lazily initializes a zero List value. +func (l *List[T]) lazyInit() { + if l.root.next == nil { + l.Init() + } +} + +// insert inserts e after at, increments l.len, and returns e. +func (l *List[T]) insert(e, at *Element[T]) *Element[T] { + e.prev = at + e.next = at.next + e.prev.next = e + e.next.prev = e + e.list = l + l.len++ + return e +} + +// insertValue is a convenience wrapper for insert(&Element{Value: v}, at). +func (l *List[T]) insertValue(v T, at *Element[T]) *Element[T] { + var e *Element[T] + if l.pool != nil { + e = l.pool.Get().(*Element[T]) + } else { + e = &Element[T]{} + } + e.Value = v + return l.insert(e, at) +} + +// remove removes e from its list, decrements l.len +func (l *List[T]) remove(e *Element[T]) { + e.prev.next = e.next + e.next.prev = e.prev + e.next = nil // avoid memory leaks + e.prev = nil // avoid memory leaks + e.list = nil + if l.pool != nil { + l.pool.Put(e) + } + l.len-- +} + +// move moves e to next to at. +func (l *List[T]) move(e, at *Element[T]) { + if e == at { + return + } + e.prev.next = e.next + e.next.prev = e.prev + + e.prev = at + e.next = at.next + e.prev.next = e + e.next.prev = e +} + +// Remove removes e from l if e is an element of list l. +// It returns the element value e.Value. +// The element must not be nil. +func (l *List[T]) Remove(e *Element[T]) T { + v := e.Value + if e.list == l { + // if e.list == l, l must have been initialized when e was inserted + // in l or l == nil (e is a zero Element) and l.remove will crash + l.remove(e) + } + return v +} + +// PushFront inserts a new element e with value v at the front of list l and returns e. +func (l *List[T]) PushFront(v T) *Element[T] { + l.lazyInit() + return l.insertValue(v, &l.root) +} + +// PushBack inserts a new element e with value v at the back of list l and returns e. +func (l *List[T]) PushBack(v T) *Element[T] { + l.lazyInit() + return l.insertValue(v, l.root.prev) +} + +// InsertBefore inserts a new element e with value v immediately before mark and returns e. +// If mark is not an element of l, the list is not modified. +// The mark must not be nil. +func (l *List[T]) InsertBefore(v T, mark *Element[T]) *Element[T] { + if mark.list != l { + return nil + } + // see comment in List.Remove about initialization of l + return l.insertValue(v, mark.prev) +} + +// InsertAfter inserts a new element e with value v immediately after mark and returns e. +// If mark is not an element of l, the list is not modified. +// The mark must not be nil. +func (l *List[T]) InsertAfter(v T, mark *Element[T]) *Element[T] { + if mark.list != l { + return nil + } + // see comment in List.Remove about initialization of l + return l.insertValue(v, mark) +} + +// MoveToFront moves element e to the front of list l. +// If e is not an element of l, the list is not modified. +// The element must not be nil. +func (l *List[T]) MoveToFront(e *Element[T]) { + if e.list != l || l.root.next == e { + return + } + // see comment in List.Remove about initialization of l + l.move(e, &l.root) +} + +// MoveToBack moves element e to the back of list l. +// If e is not an element of l, the list is not modified. +// The element must not be nil. +func (l *List[T]) MoveToBack(e *Element[T]) { + if e.list != l || l.root.prev == e { + return + } + // see comment in List.Remove about initialization of l + l.move(e, l.root.prev) +} + +// MoveBefore moves element e to its new position before mark. +// If e or mark is not an element of l, or e == mark, the list is not modified. +// The element and mark must not be nil. +func (l *List[T]) MoveBefore(e, mark *Element[T]) { + if e.list != l || e == mark || mark.list != l { + return + } + l.move(e, mark.prev) +} + +// MoveAfter moves element e to its new position after mark. +// If e or mark is not an element of l, or e == mark, the list is not modified. +// The element and mark must not be nil. +func (l *List[T]) MoveAfter(e, mark *Element[T]) { + if e.list != l || e == mark || mark.list != l { + return + } + l.move(e, mark) +} + +// PushBackList inserts a copy of another list at the back of list l. +// The lists l and other may be the same. They must not be nil. +func (l *List[T]) PushBackList(other *List[T]) { + l.lazyInit() + for i, e := other.Len(), other.Front(); i > 0; i, e = i-1, e.Next() { + l.insertValue(e.Value, l.root.prev) + } +} + +// PushFrontList inserts a copy of another list at the front of list l. +// The lists l and other may be the same. They must not be nil. +func (l *List[T]) PushFrontList(other *List[T]) { + l.lazyInit() + for i, e := other.Len(), other.Back(); i > 0; i, e = i-1, e.Prev() { + l.insertValue(e.Value, &l.root) + } +} diff --git a/third_party/quic-go/internal/utils/log.go b/third_party/quic-go/internal/utils/log.go new file mode 100644 index 0000000..d10c5ac --- /dev/null +++ b/third_party/quic-go/internal/utils/log.go @@ -0,0 +1,131 @@ +package utils + +import ( + "fmt" + "log" + "os" + "strings" + "time" +) + +// LogLevel of quic-go +type LogLevel uint8 + +const ( + // LogLevelNothing disables + LogLevelNothing LogLevel = iota + // LogLevelError enables err logs + LogLevelError + // LogLevelInfo enables info logs (e.g. packets) + LogLevelInfo + // LogLevelDebug enables debug logs (e.g. packet contents) + LogLevelDebug +) + +const logEnv = "QUIC_GO_LOG_LEVEL" + +// A Logger logs. +type Logger interface { + SetLogLevel(LogLevel) + SetLogTimeFormat(format string) + WithPrefix(prefix string) Logger + Debug() bool + + Errorf(format string, args ...any) + Infof(format string, args ...any) + Debugf(format string, args ...any) +} + +// DefaultLogger is used by quic-go for logging. +var DefaultLogger Logger + +type defaultLogger struct { + prefix string + + logLevel LogLevel + timeFormat string +} + +var _ Logger = &defaultLogger{} + +// SetLogLevel sets the log level +func (l *defaultLogger) SetLogLevel(level LogLevel) { + l.logLevel = level +} + +// SetLogTimeFormat sets the format of the timestamp +// an empty string disables the logging of timestamps +func (l *defaultLogger) SetLogTimeFormat(format string) { + log.SetFlags(0) // disable timestamp logging done by the log package + l.timeFormat = format +} + +// Debugf logs something +func (l *defaultLogger) Debugf(format string, args ...any) { + if l.logLevel == LogLevelDebug { + l.logMessage(format, args...) + } +} + +// Infof logs something +func (l *defaultLogger) Infof(format string, args ...any) { + if l.logLevel >= LogLevelInfo { + l.logMessage(format, args...) + } +} + +// Errorf logs something +func (l *defaultLogger) Errorf(format string, args ...any) { + if l.logLevel >= LogLevelError { + l.logMessage(format, args...) + } +} + +func (l *defaultLogger) logMessage(format string, args ...any) { + var pre string + + if len(l.timeFormat) > 0 { + pre = time.Now().Format(l.timeFormat) + " " + } + if len(l.prefix) > 0 { + pre += l.prefix + " " + } + log.Printf(pre+format, args...) +} + +func (l *defaultLogger) WithPrefix(prefix string) Logger { + if len(l.prefix) > 0 { + prefix = l.prefix + " " + prefix + } + return &defaultLogger{ + logLevel: l.logLevel, + timeFormat: l.timeFormat, + prefix: prefix, + } +} + +// Debug returns true if the log level is LogLevelDebug +func (l *defaultLogger) Debug() bool { + return l.logLevel == LogLevelDebug +} + +func init() { + DefaultLogger = &defaultLogger{} + DefaultLogger.SetLogLevel(readLoggingEnv()) +} + +func readLoggingEnv() LogLevel { + switch strings.ToLower(os.Getenv(logEnv)) { + case "": + return LogLevelNothing + case "debug": + return LogLevelDebug + case "info": + return LogLevelInfo + case "error": + return LogLevelError + default: + fmt.Fprintln(os.Stderr, "invalid quic-go log level, see https://github.com/apernet/quic-go/wiki/Logging") + return LogLevelNothing + } +} diff --git a/third_party/quic-go/internal/utils/log_test.go b/third_party/quic-go/internal/utils/log_test.go new file mode 100644 index 0000000..8fc7d05 --- /dev/null +++ b/third_party/quic-go/internal/utils/log_test.go @@ -0,0 +1,146 @@ +package utils + +import ( + "bytes" + "log" + "os" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestLogLevelNothing(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + DefaultLogger.SetLogLevel(LogLevelNothing) + DefaultLogger.Debugf("debug") + DefaultLogger.Infof("info") + DefaultLogger.Errorf("err") + require.Empty(t, b.String()) +} + +func TestLogLevelError(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + DefaultLogger.SetLogLevel(LogLevelError) + DefaultLogger.Debugf("debug") + DefaultLogger.Infof("info") + DefaultLogger.Errorf("err") + require.Contains(t, b.String(), "err\n") + require.NotContains(t, b.String(), "info") + require.NotContains(t, b.String(), "debug") +} + +func TestLogLevelInfo(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + DefaultLogger.SetLogLevel(LogLevelInfo) + DefaultLogger.Debugf("debug") + DefaultLogger.Infof("info") + DefaultLogger.Errorf("err") + require.Contains(t, b.String(), "err\n") + require.Contains(t, b.String(), "info\n") + require.NotContains(t, b.String(), "debug") +} + +func TestLogLevelDebug(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + require.False(t, DefaultLogger.Debug()) + DefaultLogger.SetLogLevel(LogLevelDebug) + require.True(t, DefaultLogger.Debug()) + DefaultLogger.Debugf("debug") + DefaultLogger.Infof("info") + DefaultLogger.Errorf("err") + require.Contains(t, b.String(), "err\n") + require.Contains(t, b.String(), "info\n") + require.Contains(t, b.String(), "debug\n") +} + +func TestNoTimestampWithEmptyFormat(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + DefaultLogger.SetLogLevel(LogLevelDebug) + DefaultLogger.SetLogTimeFormat("") + DefaultLogger.Debugf("debug") + require.Equal(t, "debug\n", b.String()) +} + +func TestAddTimestamp(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + format := "Jan 2, 2006" + DefaultLogger.SetLogTimeFormat(format) + DefaultLogger.SetLogLevel(LogLevelInfo) + DefaultLogger.Infof("info") + timestamp := b.String()[:b.Len()-6] + parsedTime, err := time.ParseInLocation(format, timestamp, time.Local) + require.NoError(t, err) + require.WithinDuration(t, time.Now(), parsedTime, 25*time.Hour) +} + +func TestLogAddPrefixes(t *testing.T) { + b := &bytes.Buffer{} + log.SetOutput(b) + defer log.SetOutput(os.Stdout) + defer DefaultLogger.SetLogLevel(LogLevelNothing) + + DefaultLogger.SetLogLevel(LogLevelDebug) + + // single prefix + prefixLogger := DefaultLogger.WithPrefix("prefix") + prefixLogger.Debugf("debug1") + require.Contains(t, b.String(), "prefix") + require.Contains(t, b.String(), "debug1") + + // multiple prefixes + b.Reset() + prefixLogger1 := DefaultLogger.WithPrefix("prefix1") + prefixLogger2 := prefixLogger1.WithPrefix("prefix2") + prefixLogger2.Debugf("debug2") + require.Contains(t, b.String(), "prefix1") + require.Contains(t, b.String(), "prefix2") + require.Contains(t, b.String(), "debug2") +} + +func TestLogLevelFromEnv(t *testing.T) { + testCases := []struct { + envValue string + expected LogLevel + }{ + {"DEBUG", LogLevelDebug}, + {"debug", LogLevelDebug}, + {"INFO", LogLevelInfo}, + {"ERROR", LogLevelError}, + } + + for _, tc := range testCases { + t.Setenv(logEnv, tc.envValue) + require.Equal(t, tc.expected, readLoggingEnv()) + } + + // invalid values + t.Setenv(logEnv, "") + require.Equal(t, LogLevelNothing, readLoggingEnv()) + t.Setenv(logEnv, "asdf") + require.Equal(t, LogLevelNothing, readLoggingEnv()) +} diff --git a/third_party/quic-go/internal/utils/rand.go b/third_party/quic-go/internal/utils/rand.go new file mode 100644 index 0000000..3006914 --- /dev/null +++ b/third_party/quic-go/internal/utils/rand.go @@ -0,0 +1,29 @@ +package utils + +import ( + "crypto/rand" + "encoding/binary" +) + +// Rand is a wrapper around crypto/rand that adds some convenience functions known from math/rand. +type Rand struct { + buf [4]byte +} + +func (r *Rand) Int31() int32 { + rand.Read(r.buf[:]) + return int32(binary.BigEndian.Uint32(r.buf[:]) & ^uint32(1<<31)) +} + +// copied from the standard library math/rand implementation of Int63n +func (r *Rand) Int31n(n int32) int32 { + if n&(n-1) == 0 { // n is power of two, can mask + return r.Int31() & (n - 1) + } + max := int32((1 << 31) - 1 - (1<<31)%uint32(n)) + v := r.Int31() + for v > max { + v = r.Int31() + } + return v % n +} diff --git a/third_party/quic-go/internal/utils/rand_test.go b/third_party/quic-go/internal/utils/rand_test.go new file mode 100644 index 0000000..fcef097 --- /dev/null +++ b/third_party/quic-go/internal/utils/rand_test.go @@ -0,0 +1,32 @@ +package utils + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestRandomNumbers(t *testing.T) { + const ( + num = 1000 + max = 12345678 + ) + + var values [num]int32 + var r Rand + for i := range num { + v := r.Int31n(max) + require.GreaterOrEqual(t, v, int32(0)) + require.Less(t, v, int32(max)) + values[i] = v + } + + var sum uint64 + for _, n := range values { + sum += uint64(n) + } + average := float64(sum) / num + expectedAverage := float64(max) / 2 + tolerance := float64(max) / 25 + require.InDelta(t, expectedAverage, average, tolerance) +} diff --git a/third_party/quic-go/internal/utils/ringbuffer/ringbuffer.go b/third_party/quic-go/internal/utils/ringbuffer/ringbuffer.go new file mode 100644 index 0000000..64f64ce --- /dev/null +++ b/third_party/quic-go/internal/utils/ringbuffer/ringbuffer.go @@ -0,0 +1,96 @@ +package ringbuffer + +// A RingBuffer is a ring buffer. +// It acts as a heap that doesn't cause any allocations. +type RingBuffer[T any] struct { + ring []T + headPos, tailPos int + full bool +} + +// Init preallocates a buffer with a certain size. +func (r *RingBuffer[T]) Init(size int) { + r.ring = make([]T, size) +} + +// Len returns the number of elements in the ring buffer. +func (r *RingBuffer[T]) Len() int { + if r.full { + return len(r.ring) + } + if r.tailPos >= r.headPos { + return r.tailPos - r.headPos + } + return r.tailPos - r.headPos + len(r.ring) +} + +// Empty says if the ring buffer is empty. +func (r *RingBuffer[T]) Empty() bool { + return !r.full && r.headPos == r.tailPos +} + +// PushBack adds a new element. +// If the ring buffer is full, its capacity is increased first. +func (r *RingBuffer[T]) PushBack(t T) { + if r.full || len(r.ring) == 0 { + r.grow() + } + r.ring[r.tailPos] = t + r.tailPos++ + if r.tailPos == len(r.ring) { + r.tailPos = 0 + } + if r.tailPos == r.headPos { + r.full = true + } +} + +// PopFront returns the next element. +// It must not be called when the buffer is empty, that means that +// callers might need to check if there are elements in the buffer first. +func (r *RingBuffer[T]) PopFront() T { + if r.Empty() { + panic("github.com/apernet/quic-go/internal/utils/ringbuffer: pop from an empty queue") + } + r.full = false + t := r.ring[r.headPos] + r.ring[r.headPos] = *new(T) + r.headPos++ + if r.headPos == len(r.ring) { + r.headPos = 0 + } + return t +} + +// PeekFront returns the next element. +// It must not be called when the buffer is empty, that means that +// callers might need to check if there are elements in the buffer first. +func (r *RingBuffer[T]) PeekFront() T { + if r.Empty() { + panic("github.com/apernet/quic-go/internal/utils/ringbuffer: peek from an empty queue") + } + return r.ring[r.headPos] +} + +// Grow the maximum size of the queue. +// This method assume the queue is full. +func (r *RingBuffer[T]) grow() { + oldRing := r.ring + newSize := len(oldRing) * 2 + if newSize == 0 { + newSize = 1 + } + r.ring = make([]T, newSize) + headLen := copy(r.ring, oldRing[r.headPos:]) + copy(r.ring[headLen:], oldRing[:r.headPos]) + r.headPos, r.tailPos, r.full = 0, len(oldRing), false +} + +// Clear removes all elements. +func (r *RingBuffer[T]) Clear() { + var zeroValue T + for i := range r.ring { + r.ring[i] = zeroValue + } + r.headPos, r.tailPos, r.full = 0, 0, false +} diff --git a/third_party/quic-go/internal/utils/ringbuffer/ringbuffer_bench_test.go b/third_party/quic-go/internal/utils/ringbuffer/ringbuffer_bench_test.go new file mode 100644 index 0000000..83db80a --- /dev/null +++ b/third_party/quic-go/internal/utils/ringbuffer/ringbuffer_bench_test.go @@ -0,0 +1,14 @@ +package ringbuffer + +import "testing" + +func BenchmarkRingBuffer(b *testing.B) { + r := RingBuffer[int]{} + + var val int + for b.Loop() { + r.PushBack(val) + r.PopFront() + val++ + } +} diff --git a/third_party/quic-go/internal/utils/ringbuffer/ringbuffer_test.go b/third_party/quic-go/internal/utils/ringbuffer/ringbuffer_test.go new file mode 100644 index 0000000..387d727 --- /dev/null +++ b/third_party/quic-go/internal/utils/ringbuffer/ringbuffer_test.go @@ -0,0 +1,49 @@ +package ringbuffer + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestPushPeekPop(t *testing.T) { + r := RingBuffer[int]{} + require.Equal(t, 0, len(r.ring)) + require.Panics(t, func() { r.PopFront() }) + r.PushBack(1) + r.PushBack(2) + r.PushBack(3) + require.Equal(t, 1, r.PeekFront()) + require.Equal(t, 1, r.PeekFront()) + require.Equal(t, 1, r.PopFront()) + require.Equal(t, 2, r.PeekFront()) + require.Equal(t, 2, r.PopFront()) + r.PushBack(4) + r.PushBack(5) + require.Equal(t, 3, r.Len()) + r.PushBack(6) + require.Equal(t, 4, r.Len()) + require.Equal(t, 3, r.PopFront()) + require.Equal(t, 4, r.PopFront()) + require.Equal(t, 5, r.PopFront()) + require.Equal(t, 6, r.PopFront()) +} + +func TestPanicOnEmptyBuffer(t *testing.T) { + r := RingBuffer[string]{} + require.True(t, r.Empty()) + require.Zero(t, r.Len()) + require.Panics(t, func() { r.PeekFront() }) + require.Panics(t, func() { r.PopFront() }) +} + +func TestClear(t *testing.T) { + r := RingBuffer[int]{} + r.Init(2) + r.PushBack(1) + r.PushBack(2) + require.True(t, r.full) + r.Clear() + require.False(t, r.full) + require.Equal(t, 0, r.Len()) +} diff --git a/third_party/quic-go/internal/utils/rtt_stats.go b/third_party/quic-go/internal/utils/rtt_stats.go new file mode 100644 index 0000000..4ea7b3b --- /dev/null +++ b/third_party/quic-go/internal/utils/rtt_stats.go @@ -0,0 +1,159 @@ +package utils + +import ( + "sync/atomic" + "time" + + "github.com/apernet/quic-go/internal/protocol" +) + +const ( + rttAlpha = 0.125 + oneMinusAlpha = 1 - rttAlpha + rttBeta = 0.25 + oneMinusBeta = 1 - rttBeta +) + +// The default RTT used before an RTT sample is taken +const DefaultInitialRTT = 100 * time.Millisecond + +// RTTStats provides round-trip statistics +type RTTStats struct { + hasMeasurement bool + + minRTT atomic.Int64 // nanoseconds + latestRTT atomic.Int64 // nanoseconds + smoothedRTT atomic.Int64 // nanoseconds + meanDeviation atomic.Int64 // nanoseconds + + maxAckDelay atomic.Int64 // nanoseconds +} + +func NewRTTStats() *RTTStats { + var rttStats RTTStats + rttStats.minRTT.Store(DefaultInitialRTT.Nanoseconds()) + rttStats.latestRTT.Store(DefaultInitialRTT.Nanoseconds()) + rttStats.smoothedRTT.Store(DefaultInitialRTT.Nanoseconds()) + return &rttStats +} + +// MinRTT Returns the minRTT for the entire connection. +// May return Zero if no valid updates have occurred. +func (r *RTTStats) MinRTT() time.Duration { + return time.Duration(r.minRTT.Load()) +} + +// LatestRTT returns the most recent rtt measurement. +// May return Zero if no valid updates have occurred. +func (r *RTTStats) LatestRTT() time.Duration { + return time.Duration(r.latestRTT.Load()) +} + +// SmoothedRTT returns the smoothed RTT for the connection. +// May return Zero if no valid updates have occurred. +func (r *RTTStats) SmoothedRTT() time.Duration { + return time.Duration(r.smoothedRTT.Load()) +} + +// MeanDeviation gets the mean deviation +func (r *RTTStats) MeanDeviation() time.Duration { + return time.Duration(r.meanDeviation.Load()) +} + +// MaxAckDelay gets the max_ack_delay advertised by the peer +func (r *RTTStats) MaxAckDelay() time.Duration { + return time.Duration(r.maxAckDelay.Load()) +} + +// PTO gets the probe timeout duration. +func (r *RTTStats) PTO(includeMaxAckDelay bool) time.Duration { + if !r.hasMeasurement { + return 2 * DefaultInitialRTT + } + pto := r.SmoothedRTT() + max(4*r.MeanDeviation(), protocol.TimerGranularity) + if includeMaxAckDelay { + pto += r.MaxAckDelay() + } + return pto +} + +// UpdateRTT updates the RTT based on a new sample. +func (r *RTTStats) UpdateRTT(sendDelta, ackDelay time.Duration) { + if sendDelta <= 0 { + return + } + + // Update r.minRTT first. r.minRTT does not use an rttSample corrected for + // ackDelay but the raw observed sendDelta, since poor clock granularity at + // the client may cause a high ackDelay to result in underestimation of the + // r.minRTT. + minRTT := time.Duration(r.minRTT.Load()) + if !r.hasMeasurement || minRTT > sendDelta { + minRTT = sendDelta + r.minRTT.Store(sendDelta.Nanoseconds()) + } + + // Correct for ackDelay if information received from the peer results in a + // an RTT sample at least as large as minRTT. Otherwise, only use the + // sendDelta. + sample := sendDelta + if sample-minRTT >= ackDelay { + sample -= ackDelay + } + r.latestRTT.Store(sample.Nanoseconds()) + // First time call. + if !r.hasMeasurement { + r.hasMeasurement = true + r.smoothedRTT.Store(sample.Nanoseconds()) + r.meanDeviation.Store(sample.Nanoseconds() / 2) + } else { + smoothedRTT := r.SmoothedRTT() + meanDev := time.Duration(oneMinusBeta*float32(r.MeanDeviation()/time.Microsecond)+rttBeta*float32((smoothedRTT-sample).Abs()/time.Microsecond)) * time.Microsecond + newSmoothedRTT := time.Duration((float32(smoothedRTT/time.Microsecond)*oneMinusAlpha)+(float32(sample/time.Microsecond)*rttAlpha)) * time.Microsecond + r.meanDeviation.Store(meanDev.Nanoseconds()) + r.smoothedRTT.Store(newSmoothedRTT.Nanoseconds()) + } +} + +func (r *RTTStats) HasMeasurement() bool { + return r.hasMeasurement +} + +// SetMaxAckDelay sets the max_ack_delay +func (r *RTTStats) SetMaxAckDelay(mad time.Duration) { + r.maxAckDelay.Store(int64(mad)) +} + +// SetInitialRTT sets the initial RTT. +// It is used during handshake when restoring the RTT stats from the token. +func (r *RTTStats) SetInitialRTT(t time.Duration) { + // On the server side, by the time we get to process the session ticket, + // we might already have obtained an RTT measurement. + // This can happen if we received the ClientHello in multiple pieces, and one of those pieces was lost. + // Discard the restored value. A fresh measurement is always better. + if r.hasMeasurement { + return + } + r.smoothedRTT.Store(int64(t)) + r.latestRTT.Store(int64(t)) +} + +func (r *RTTStats) ResetForPathMigration() { + r.hasMeasurement = false + r.minRTT.Store(DefaultInitialRTT.Nanoseconds()) + r.latestRTT.Store(DefaultInitialRTT.Nanoseconds()) + r.smoothedRTT.Store(DefaultInitialRTT.Nanoseconds()) + r.meanDeviation.Store(0) + // max_ack_delay remains valid +} + +func (r *RTTStats) Clone() *RTTStats { + out := &RTTStats{} + out.hasMeasurement = r.hasMeasurement + out.minRTT.Store(r.minRTT.Load()) + out.latestRTT.Store(r.latestRTT.Load()) + out.smoothedRTT.Store(r.smoothedRTT.Load()) + out.meanDeviation.Store(r.meanDeviation.Load()) + out.maxAckDelay.Store(r.maxAckDelay.Load()) + return out +} diff --git a/third_party/quic-go/internal/utils/rtt_stats_test.go b/third_party/quic-go/internal/utils/rtt_stats_test.go new file mode 100644 index 0000000..6277db3 --- /dev/null +++ b/third_party/quic-go/internal/utils/rtt_stats_test.go @@ -0,0 +1,146 @@ +package utils + +import ( + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/stretchr/testify/require" +) + +func TestRTTStatsDefaults(t *testing.T) { + rttStats := NewRTTStats() + require.False(t, rttStats.HasMeasurement()) + require.Equal(t, DefaultInitialRTT, rttStats.MinRTT()) + require.Equal(t, DefaultInitialRTT, rttStats.SmoothedRTT()) +} + +func TestRTTStatsSmoothedRTT(t *testing.T) { + rttStats := NewRTTStats() + require.False(t, rttStats.HasMeasurement()) + // verify that ack_delay is ignored in the first measurement + rttStats.UpdateRTT(300*time.Millisecond, 100*time.Millisecond) + require.True(t, rttStats.HasMeasurement()) + require.Equal(t, 300*time.Millisecond, rttStats.LatestRTT()) + require.Equal(t, 300*time.Millisecond, rttStats.SmoothedRTT()) + // verify that smoothed RTT includes max ack delay if it's reasonable + rttStats.UpdateRTT(350*time.Millisecond, 50*time.Millisecond) + require.Equal(t, 300*time.Millisecond, rttStats.LatestRTT()) + require.Equal(t, 300*time.Millisecond, rttStats.SmoothedRTT()) + // verify that large erroneous ack_delay does not change smoothed RTT + rttStats.UpdateRTT(200*time.Millisecond, 300*time.Millisecond) + require.Equal(t, 200*time.Millisecond, rttStats.LatestRTT()) + require.Equal(t, 287500*time.Microsecond, rttStats.SmoothedRTT()) +} + +func TestRTTStatsMinRTT(t *testing.T) { + rttStats := NewRTTStats() + rttStats.UpdateRTT(200*time.Millisecond, 0) + require.Equal(t, 200*time.Millisecond, rttStats.MinRTT()) + rttStats.UpdateRTT(10*time.Millisecond, 0) + require.Equal(t, 10*time.Millisecond, rttStats.MinRTT()) + rttStats.UpdateRTT(50*time.Millisecond, 0) + require.Equal(t, 10*time.Millisecond, rttStats.MinRTT()) + rttStats.UpdateRTT(50*time.Millisecond, 0) + require.Equal(t, 10*time.Millisecond, rttStats.MinRTT()) + rttStats.UpdateRTT(50*time.Millisecond, 0) + require.Equal(t, 10*time.Millisecond, rttStats.MinRTT()) + // verify that ack_delay does not go into recording of MinRTT + rttStats.UpdateRTT(7*time.Millisecond, 2*time.Millisecond) + require.Equal(t, 7*time.Millisecond, rttStats.MinRTT()) +} + +func TestRTTStatsMaxAckDelay(t *testing.T) { + rttStats := NewRTTStats() + rttStats.SetMaxAckDelay(42 * time.Minute) + require.Equal(t, 42*time.Minute, rttStats.MaxAckDelay()) +} + +func TestRTTStatsComputePTO(t *testing.T) { + const ( + maxAckDelay = 42 * time.Minute + rtt = time.Second + ) + rttStats := NewRTTStats() + rttStats.SetMaxAckDelay(maxAckDelay) + rttStats.UpdateRTT(rtt, 0) + require.Equal(t, rtt, rttStats.SmoothedRTT()) + require.Equal(t, rtt/2, rttStats.MeanDeviation()) + require.Equal(t, rtt+4*(rtt/2), rttStats.PTO(false)) + require.Equal(t, rtt+4*(rtt/2)+maxAckDelay, rttStats.PTO(true)) +} + +func TestRTTStatsPTOWithShortRTT(t *testing.T) { + const rtt = time.Microsecond + rttStats := NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + require.Equal(t, rtt+protocol.TimerGranularity, rttStats.PTO(true)) +} + +func TestRTTStatsUpdateWithBadSendDeltas(t *testing.T) { + rttStats := NewRTTStats() + const initialRtt = 10 * time.Millisecond + rttStats.UpdateRTT(initialRtt, 0) + require.Equal(t, initialRtt, rttStats.MinRTT()) + require.Equal(t, initialRtt, rttStats.SmoothedRTT()) + + badSendDeltas := []time.Duration{ + 0, + -1000 * time.Microsecond, + } + + for _, badSendDelta := range badSendDeltas { + rttStats.UpdateRTT(badSendDelta, 0) + require.Equal(t, initialRtt, rttStats.MinRTT()) + require.Equal(t, initialRtt, rttStats.SmoothedRTT()) + } +} + +func TestRTTStatsRestore(t *testing.T) { + rttStats := NewRTTStats() + rttStats.SetInitialRTT(10 * time.Second) + require.Equal(t, 10*time.Second, rttStats.LatestRTT()) + require.Equal(t, 10*time.Second, rttStats.SmoothedRTT()) + require.Zero(t, rttStats.MeanDeviation()) + // update the RTT and make sure that the initial value is immediately forgotten + rttStats.UpdateRTT(200*time.Millisecond, 0) + require.Equal(t, 200*time.Millisecond, rttStats.LatestRTT()) + require.Equal(t, 200*time.Millisecond, rttStats.SmoothedRTT()) + require.Equal(t, 100*time.Millisecond, rttStats.MeanDeviation()) +} + +func TestRTTMeasurementAfterRestore(t *testing.T) { + rttStats := NewRTTStats() + const rtt = 10 * time.Millisecond + rttStats.UpdateRTT(rtt, 0) + require.Equal(t, rtt, rttStats.LatestRTT()) + require.Equal(t, rtt, rttStats.SmoothedRTT()) + rttStats.SetInitialRTT(time.Minute) + require.Equal(t, rtt, rttStats.LatestRTT()) + require.Equal(t, rtt, rttStats.SmoothedRTT()) +} + +func TestRTTStatsResetForPathMigration(t *testing.T) { + rttStats := NewRTTStats() + rttStats.SetMaxAckDelay(42 * time.Millisecond) + rttStats.UpdateRTT(time.Second, 0) + rttStats.UpdateRTT(10*time.Second, 0) + require.True(t, rttStats.HasMeasurement()) + require.Equal(t, time.Second, rttStats.MinRTT()) + require.Equal(t, 10*time.Second, rttStats.LatestRTT()) + require.NotZero(t, rttStats.SmoothedRTT()) + + rttStats.ResetForPathMigration() + require.False(t, rttStats.HasMeasurement()) + require.Equal(t, DefaultInitialRTT, rttStats.MinRTT()) + require.Equal(t, DefaultInitialRTT, rttStats.LatestRTT()) + require.Equal(t, DefaultInitialRTT, rttStats.SmoothedRTT()) + require.Equal(t, 2*DefaultInitialRTT, rttStats.PTO(false)) + // make sure that max_ack_delay was not reset + require.Equal(t, 42*time.Millisecond, rttStats.MaxAckDelay()) + + rttStats.UpdateRTT(10*time.Millisecond, 0) + require.True(t, rttStats.HasMeasurement()) + require.Equal(t, 10*time.Millisecond, rttStats.SmoothedRTT()) + require.Equal(t, 10*time.Millisecond, rttStats.LatestRTT()) +} diff --git a/third_party/quic-go/internal/utils/streamframe_interval.go b/third_party/quic-go/internal/utils/streamframe_interval.go new file mode 100644 index 0000000..e63ea2f --- /dev/null +++ b/third_party/quic-go/internal/utils/streamframe_interval.go @@ -0,0 +1,45 @@ +package utils + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/protocol" +) + +// ByteInterval is an interval from one ByteCount to the other +type ByteInterval struct { + Start protocol.ByteCount + End protocol.ByteCount +} + +func (i ByteInterval) Comp(v ByteInterval) int8 { + if i.Start < v.Start { + return -1 + } + if i.Start > v.Start { + return 1 + } + if i.End < v.End { + return -1 + } + if i.End > v.End { + return 1 + } + return 0 +} + +func (i ByteInterval) Match(n ByteInterval) int8 { + // check if there is an overlap + if i.Start <= n.End && i.End >= n.Start { + return 0 + } + if i.Start > n.End { + return 1 + } else { + return -1 + } +} + +func (i ByteInterval) String() string { + return fmt.Sprintf("[%d, %d]", i.Start, i.End) +} diff --git a/third_party/quic-go/internal/utils/tree/tree.go b/third_party/quic-go/internal/utils/tree/tree.go new file mode 100644 index 0000000..0d44a65 --- /dev/null +++ b/third_party/quic-go/internal/utils/tree/tree.go @@ -0,0 +1,503 @@ +// Originated from https://github.com/ross-oreto/go-tree/blob/master/btree.go with the following changes: +// 1. Genericized the code +// 2. Added Match function for our frame sorter use case +// 3. Fixed a bug in deleteNode where in some cases the deleted flag was not set to true + +/* +Copyright (c) 2017 Ross Oreto + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +*/ + +package tree + +import ( + "fmt" + "slices" +) + +type Val[T any] interface { + Comp(val T) int8 // returns 1 if > val, -1 if < val, 0 if equals to val + Match(cond T) int8 // returns 1 if > cond, -1 if < cond, 0 if matches cond +} + +// Btree represents an AVL tree +type Btree[T Val[T]] struct { + root *Node[T] + values []T + len int +} + +// Node represents a node in the tree with a value, left and right children, and a height/balance of the node. +type Node[T Val[T]] struct { + Value T + left, right *Node[T] + height int8 +} + +// New returns a new btree +func New[T Val[T]]() *Btree[T] { return new(Btree[T]).Init() } + +// Init initializes all values/clears the tree and returns the tree pointer +func (t *Btree[T]) Init() *Btree[T] { + t.root = nil + t.values = nil + t.len = 0 + return t +} + +// String returns a string representation of the tree values +func (t *Btree[T]) String() string { + return fmt.Sprint(t.Values()) +} + +// Empty returns true if the tree is empty +func (t *Btree[T]) Empty() bool { + return t.root == nil +} + +// NotEmpty returns true if the tree is not empty +func (t *Btree[T]) NotEmpty() bool { + return t.root != nil +} + +// Insert inserts a new value into the tree and returns the tree pointer +func (t *Btree[T]) Insert(value T) *Btree[T] { + added := false + t.root = insert(t.root, value, &added) + if added { + t.len++ + } + t.values = nil + return t +} + +func insert[T Val[T]](n *Node[T], value T, added *bool) *Node[T] { + if n == nil { + *added = true + return (&Node[T]{Value: value}).Init() + } + c := value.Comp(n.Value) + if c > 0 { + n.right = insert(n.right, value, added) + } else if c < 0 { + n.left = insert(n.left, value, added) + } else { + n.Value = value + *added = false + return n + } + + n.height = n.maxHeight() + 1 + c = balance(n) + + if c > 1 { + c = value.Comp(n.left.Value) + if c < 0 { + return n.rotateRight() + } else if c > 0 { + n.left = n.left.rotateLeft() + return n.rotateRight() + } + } else if c < -1 { + c = value.Comp(n.right.Value) + if c > 0 { + return n.rotateLeft() + } else if c < 0 { + n.right = n.right.rotateRight() + return n.rotateLeft() + } + } + return n +} + +// InsertAll inserts all the values into the tree and returns the tree pointer +func (t *Btree[T]) InsertAll(values []T) *Btree[T] { + for _, v := range values { + t.Insert(v) + } + return t +} + +// Contains returns true if the tree contains the specified value +func (t *Btree[T]) Contains(value T) bool { + return t.Get(value) != nil +} + +// ContainsAny returns true if the tree contains any of the values +func (t *Btree[T]) ContainsAny(values []T) bool { + return slices.ContainsFunc(values, t.Contains) +} + +// ContainsAll returns true if the tree contains all of the values +func (t *Btree[T]) ContainsAll(values []T) bool { + for _, v := range values { + if !t.Contains(v) { + return false + } + } + return true +} + +// Get returns the node value associated with the search value +func (t *Btree[T]) Get(value T) *T { + var node *Node[T] + if t.root != nil { + node = t.root.get(value) + } + if node != nil { + return &node.Value + } + return nil +} + +func (t *Btree[T]) Match(cond T) []T { + var matches []T + if t.root != nil { + t.root.match(cond, &matches) + } + return matches +} + +// Len return the number of nodes in the tree +func (t *Btree[T]) Len() int { + return t.len +} + +// Head returns the first value in the tree +func (t *Btree[T]) Head() *T { + if t.root == nil { + return nil + } + beginning := t.root + for beginning.left != nil { + beginning = beginning.left + } + if beginning == nil { + for beginning.right != nil { + beginning = beginning.right + } + } + if beginning != nil { + return &beginning.Value + } + return nil +} + +// Tail returns the last value in the tree +func (t *Btree[T]) Tail() *T { + if t.root == nil { + return nil + } + beginning := t.root + for beginning.right != nil { + beginning = beginning.right + } + if beginning == nil { + for beginning.left != nil { + beginning = beginning.left + } + } + if beginning != nil { + return &beginning.Value + } + return nil +} + +// Values returns a slice of all the values in tree in order +func (t *Btree[T]) Values() []T { + if t.values == nil { + t.values = make([]T, t.len) + t.Ascend(func(n *Node[T], i int) bool { + t.values[i] = n.Value + return true + }) + } + return t.values +} + +// Delete deletes the node from the tree associated with the search value +func (t *Btree[T]) Delete(value T) *Btree[T] { + deleted := false + t.root = deleteNode(t.root, value, &deleted) + if deleted { + t.len-- + } + t.values = nil + return t +} + +// DeleteAll deletes the nodes from the tree associated with the search values +func (t *Btree[T]) DeleteAll(values []T) *Btree[T] { + for _, v := range values { + t.Delete(v) + } + return t +} + +func deleteNode[T Val[T]](n *Node[T], value T, deleted *bool) *Node[T] { + if n == nil { + return n + } + + c := value.Comp(n.Value) + + if c < 0 { + n.left = deleteNode(n.left, value, deleted) + } else if c > 0 { + n.right = deleteNode(n.right, value, deleted) + } else { + if n.left == nil { + t := n.right + n.Init() + *deleted = true + return t + } else if n.right == nil { + t := n.left + n.Init() + *deleted = true + return t + } + t := n.right.min() + n.Value = t.Value + n.right = deleteNode(n.right, t.Value, deleted) + *deleted = true + } + + // re-balance + if n == nil { + return n + } + n.height = n.maxHeight() + 1 + bal := balance(n) + if bal > 1 { + if balance(n.left) >= 0 { + return n.rotateRight() + } + n.left = n.left.rotateLeft() + return n.rotateRight() + } else if bal < -1 { + if balance(n.right) <= 0 { + return n.rotateLeft() + } + n.right = n.right.rotateRight() + return n.rotateLeft() + } + + return n +} + +// Pop deletes the last node from the tree and returns its value +func (t *Btree[T]) Pop() *T { + value := t.Tail() + if value != nil { + t.Delete(*value) + } + return value +} + +// Pull deletes the first node from the tree and returns its value +func (t *Btree[T]) Pull() *T { + value := t.Head() + if value != nil { + t.Delete(*value) + } + return value +} + +// NodeIterator expresses the iterator function used for traversals +type NodeIterator[T Val[T]] func(n *Node[T], i int) bool + +// Ascend performs an ascending order traversal of the tree calling the iterator function on each node +// the iterator will continue as long as the NodeIterator returns true +func (t *Btree[T]) Ascend(iterator NodeIterator[T]) { + var i int + if t.root != nil { + t.root.iterate(iterator, &i, true) + } +} + +// Descend performs a descending order traversal of the tree using the iterator +// the iterator will continue as long as the NodeIterator returns true +func (t *Btree[T]) Descend(iterator NodeIterator[T]) { + var i int + if t.root != nil { + t.root.rIterate(iterator, &i, true) + } +} + +// Debug prints out useful debug information about the tree for debugging purposes +func (t *Btree[T]) Debug() { + fmt.Println("----------------------------------------------------------------------------------------------") + if t.Empty() { + fmt.Println("tree is empty") + } else { + fmt.Println(t.Len(), "elements") + } + + t.Ascend(func(n *Node[T], i int) bool { + if t.root.Value.Comp(n.Value) == 0 { + fmt.Print("ROOT ** ") + } + n.Debug() + return true + }) + fmt.Println("----------------------------------------------------------------------------------------------") +} + +// Init initializes the values of the node or clears the node and returns the node pointer +func (n *Node[T]) Init() *Node[T] { + n.height = 1 + n.left = nil + n.right = nil + return n +} + +// String returns a string representing the node +func (n *Node[T]) String() string { + return fmt.Sprint(n.Value) +} + +// Debug prints out useful debug information about the tree node for debugging purposes +func (n *Node[T]) Debug() { + var children string + if n.left == nil && n.right == nil { + children = "no children |" + } else if n.left != nil && n.right != nil { + children = fmt.Sprint("left child:", n.left.String(), " right child:", n.right.String()) + } else if n.right != nil { + children = fmt.Sprint("right child:", n.right.String()) + } else { + children = fmt.Sprint("left child:", n.left.String()) + } + + fmt.Println(n.String(), "|", "height", n.height, "|", "balance", balance(n), "|", children) +} + +func height[T Val[T]](n *Node[T]) int8 { + if n != nil { + return n.height + } + return 0 +} + +func balance[T Val[T]](n *Node[T]) int8 { + if n == nil { + return 0 + } + return height(n.left) - height(n.right) +} + +func (n *Node[T]) get(val T) *Node[T] { + var node *Node[T] + c := val.Comp(n.Value) + if c < 0 { + if n.left != nil { + node = n.left.get(val) + } + } else if c > 0 { + if n.right != nil { + node = n.right.get(val) + } + } else { + node = n + } + return node +} + +func (n *Node[T]) match(cond T, results *[]T) { + c := n.Value.Match(cond) + if c > 0 { + if n.left != nil { + n.left.match(cond, results) + } + } else if c < 0 { + if n.right != nil { + n.right.match(cond, results) + } + } else { + // other matching nodes could be on both sides + if n.left != nil { + n.left.match(cond, results) + } + *results = append(*results, n.Value) + if n.right != nil { + n.right.match(cond, results) + } + } +} + +func (n *Node[T]) rotateRight() *Node[T] { + l := n.left + // Rotation + l.right, n.left = n, l.right + + // update heights + n.height = n.maxHeight() + 1 + l.height = l.maxHeight() + 1 + + return l +} + +func (n *Node[T]) rotateLeft() *Node[T] { + r := n.right + // Rotation + r.left, n.right = n, r.left + + // update heights + n.height = n.maxHeight() + 1 + r.height = r.maxHeight() + 1 + + return r +} + +func (n *Node[T]) iterate(iterator NodeIterator[T], i *int, cont bool) { + if n != nil && cont { + n.left.iterate(iterator, i, cont) + cont = iterator(n, *i) + *i++ + n.right.iterate(iterator, i, cont) + } +} + +func (n *Node[T]) rIterate(iterator NodeIterator[T], i *int, cont bool) { + if n != nil && cont { + n.right.iterate(iterator, i, cont) + cont = iterator(n, *i) + *i++ + n.left.iterate(iterator, i, cont) + } +} + +func (n *Node[T]) min() *Node[T] { + current := n + for current.left != nil { + current = current.left + } + return current +} + +func (n *Node[T]) maxHeight() int8 { + rh := height(n.right) + lh := height(n.left) + if rh > lh { + return rh + } + return lh +} diff --git a/third_party/quic-go/internal/utils/tree/tree_match_test.go b/third_party/quic-go/internal/utils/tree/tree_match_test.go new file mode 100644 index 0000000..bb6c5b1 --- /dev/null +++ b/third_party/quic-go/internal/utils/tree/tree_match_test.go @@ -0,0 +1,95 @@ +package tree + +import ( + "testing" +) + +type interval struct { + start, end int +} + +func (i interval) Comp(ot interval) int8 { + if i.start < ot.start { + return -1 + } + if i.start > ot.start { + return 1 + } + if i.end < ot.end { + return -1 + } + if i.end > ot.end { + return 1 + } + return 0 +} + +func (i interval) Match(ot interval) int8 { + // Check for overlap + if i.start <= ot.end && i.end >= ot.start { + return 0 + } + if i.start > ot.end { + return 1 + } else { + return -1 + } +} + +func TestBtree(t *testing.T) { + values := []interval{ + {start: 9, end: 10}, + {start: 3, end: 4}, + {start: 1, end: 2}, + {start: 5, end: 6}, + {start: 7, end: 8}, + {start: 20, end: 100}, + {start: 11, end: 12}, + } + btree := New[interval]() + btree.InsertAll(values) + + expect, actual := len(values), btree.Len() + if actual != expect { + t.Error("length should equal", expect, "actual", actual) + } + + rs := btree.Match(interval{start: 1, end: 6}) + if len(rs) != 3 { + t.Errorf("expected 3 results, got %d", len(rs)) + } + if rs[0].start != 1 || rs[0].end != 2 { + t.Errorf("expected result 1 to be [1, 2], got %v", rs[0]) + } + if rs[1].start != 3 || rs[1].end != 4 { + t.Errorf("expected result 2 to be [3, 4], got %v", rs[1]) + } + if rs[2].start != 5 || rs[2].end != 6 { + t.Errorf("expected result 3 to be [5, 6], got %v", rs[2]) + } + + btree.Delete(interval{start: 5, end: 6}) + + rs = btree.Match(interval{start: 1, end: 6}) + if len(rs) != 2 { + t.Errorf("expected 2 results, got %d", len(rs)) + } + if rs[0].start != 1 || rs[0].end != 2 { + t.Errorf("expected result 1 to be [1, 2], got %v", rs[0]) + } + if rs[1].start != 3 || rs[1].end != 4 { + t.Errorf("expected result 2 to be [3, 4], got %v", rs[1]) + } + + btree.Delete(interval{start: 11, end: 12}) + + rs = btree.Match(interval{start: 12, end: 19}) + if len(rs) != 0 { + t.Errorf("expected 0 results, got %d", len(rs)) + } + + expect, actual = len(values)-2, btree.Len() + if actual != expect { + t.Error("length should equal", expect, "actual", actual) + } +} diff --git a/third_party/quic-go/internal/utils/tree/tree_test.go b/third_party/quic-go/internal/utils/tree/tree_test.go new file mode 100644 index 0000000..547bf4d --- /dev/null +++ b/third_party/quic-go/internal/utils/tree/tree_test.go @@ -0,0 +1,254 @@ +package tree + +import ( + "reflect" + "testing" +) + +type IntVal int + +func (i IntVal) Comp(v IntVal) int8 { + if i > v { + return 1 + } else if i < v { + return -1 + } else { + return 0 + } +} + +func (i IntVal) Match(v IntVal) int8 { + // Unused + return 0 +} + +type StringVal string + +func (i StringVal) Comp(v StringVal) int8 { + if i > v { + return 1 + } else if i < v { + return -1 + } else { + return 0 + } +} + +func (i StringVal) Match(v StringVal) int8 { + // Unused + return 0 +} + +func btreeInOrder(n int) *Btree[IntVal] { + btree := New[IntVal]() + for i := 1; i <= n; i++ { + btree.Insert(IntVal(i)) + } + return btree +} + +func btreeFixed[T Val[T]](values []T) *Btree[T] { + btree := New[T]() + btree.InsertAll(values) + return btree +} + +func TestBtree_Get(t *testing.T) { + values := []IntVal{9, 4, 2, 6, 8, 0, 3, 1, 7, 5} + btree := btreeFixed[IntVal](values).InsertAll(values) + + expect, actual := len(values), btree.Len() + if actual != expect { + t.Error("length should equal", expect, "actual", actual) + } + + expect2 := IntVal(2) + if btree.Get(expect2) == nil || *btree.Get(expect2) != expect2 { + t.Error("value should equal", expect2) + } +} + +func TestBtreeString_Get(t *testing.T) { + tree := New[StringVal]() + tree.Insert("Oreto").Insert("Michael").Insert("Ross") + + expect := StringVal("Ross") + if tree.Get(expect) == nil || *tree.Get(expect) != expect { + t.Error("value should equal", expect) + } +} + +func TestBtree_Contains(t *testing.T) { + btree := btreeInOrder(1000) + + test := IntVal(1) + if !btree.Contains(test) { + t.Error("tree should contain", test) + } + + test2 := []IntVal{1, 2, 3, 4} + if !btree.ContainsAll(test2) { + t.Error("tree should contain", test2) + } + + test2 = []IntVal{5} + if !btree.ContainsAny(test2) { + t.Error("tree should contain", test2) + } + + test2 = []IntVal{5000, 2000} + if btree.ContainsAny(test2) { + t.Error("tree should not contain any", test2) + } +} + +func TestBtree_String(t *testing.T) { + btree := btreeFixed[IntVal]([]IntVal{1, 2, 3, 4, 5, 6}) + s1 := btree.String() + s2 := "[1 2 3 4 5 6]" + if s1 != s2 { + t.Error(s1, "tree string representation should equal", s2) + } +} + +func TestBtree_Values(t *testing.T) { + const capacity = 3 + btree := btreeFixed[IntVal]([]IntVal{1, 2}) + + b := btree.Values() + c := []IntVal{1, 2} + if !reflect.DeepEqual(c, b) { + t.Error(c, "should equal", b) + } + btree.Insert(IntVal(3)) + + desc := [capacity]IntVal{} + btree.Descend(func(n *Node[IntVal], i int) bool { + desc[i] = n.Value + return true + }) + d := [capacity]IntVal{3, 2, 1} + if !reflect.DeepEqual(desc, d) { + t.Error(desc, "should equal", d) + } + + e := []IntVal{1, 2, 3} + for i, v := range btree.Values() { + if e[i] != v { + t.Error(e[i], "should equal", v) + } + } +} + +func TestBtree_Delete(t *testing.T) { + test := []IntVal{1, 2, 3} + btree := btreeFixed(test) + + btree.DeleteAll(test) + + if !btree.Empty() { + t.Error("tree should be empty") + } + + btree = btreeFixed(test) + pop := btree.Pop() + if pop == nil || *pop != IntVal(3) { + t.Error(pop, "should be 3") + } + pull := btree.Pull() + if pull == nil || *pull != IntVal(1) { + t.Error(pop, "should be 3") + } + if !btree.Delete(*btree.Pop()).Empty() { + t.Error("tree should be empty") + } + btree.Pop() + btree.Pull() +} + +func TestBtree_HeadTail(t *testing.T) { + btree := btreeFixed[IntVal]([]IntVal{1, 2, 3}) + if btree.Head() == nil || *btree.Head() != IntVal(1) { + t.Error("head element should be 1") + } + if btree.Tail() == nil || *btree.Tail() != IntVal(3) { + t.Error("head element should be 3") + } + btree.Init() + if btree.Head() != nil { + t.Error("head element should be nil") + } +} + +type TestKey1 struct { + Name string +} + +func (testkey TestKey1) Comp(tk TestKey1) int8 { + var c int8 + if testkey.Name > tk.Name { + c = 1 + } else if testkey.Name < tk.Name { + c = -1 + } + return c +} + +func (testkey TestKey1) Match(tk TestKey1) int8 { + // Unused + return 0 +} + +func TestBtree_CustomKey(t *testing.T) { + btree := New[TestKey1]() + btree.InsertAll([]TestKey1{ + {Name: "Ross"}, + {Name: "Michael"}, + {Name: "Angelo"}, + {Name: "Jason"}, + }) + + rootName := btree.root.Value.Name + if btree.root.Value.Name != "Michael" { + t.Error(rootName, "should equal Michael") + } + btree.Init() + btree.InsertAll([]TestKey1{ + {Name: "Ross"}, + {Name: "Michael"}, + {Name: "Angelo"}, + {Name: "Jason"}, + }) + btree.Debug() + s := btree.String() + test := "[{Angelo} {Jason} {Michael} {Ross}]" + if s != test { + t.Error(s, "should equal", test) + } + + btree.Delete(TestKey1{Name: "Michael"}) + if btree.Len() != 3 { + t.Error("tree length should be 3") + } + test = "Jason" + if btree.root.Value.Name != test { + t.Error(btree.root.Value, "root of the tree should be", test) + } + for !btree.Empty() { + btree.Delete(btree.root.Value) + } + btree.Debug() +} + +func TestBtree_Duplicates(t *testing.T) { + btree := New[IntVal]() + btree.InsertAll([]IntVal{ + 0, 2, 5, 10, 15, 20, 12, 14, + 13, 25, 0, 2, 5, 10, 15, 20, 12, 14, 13, 25, + }) + test := 10 + length := btree.Len() + if length != test { + t.Error(length, "tree length should be", test) + } +} diff --git a/third_party/quic-go/internal/wire/ack_frame.go b/third_party/quic-go/internal/wire/ack_frame.go new file mode 100644 index 0000000..6e53f06 --- /dev/null +++ b/third_party/quic-go/internal/wire/ack_frame.go @@ -0,0 +1,298 @@ +package wire + +import ( + "errors" + "math" + "sort" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +var errInvalidAckRanges = errors.New("AckFrame: ACK frame contains invalid ACK ranges") + +// An AckFrame is an ACK frame +type AckFrame struct { + AckRanges []AckRange // has to be ordered. The highest ACK range goes first, the lowest ACK range goes last + DelayTime time.Duration + + ECT0, ECT1, ECNCE uint64 +} + +// parseAckFrame reads an ACK frame +func parseAckFrame(frame *AckFrame, b []byte, typ FrameType, ackDelayExponent uint8, _ protocol.Version) (int, error) { + startLen := len(b) + ecn := typ == FrameTypeAckECN + + la, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + largestAcked := protocol.PacketNumber(la) + delay, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + + delayTime := time.Duration(delay*1< largestAcked { + return 0, errors.New("invalid first ACK range") + } + smallest := largestAcked - ackBlock + frame.AckRanges = append(frame.AckRanges, AckRange{Smallest: smallest, Largest: largestAcked}) + + // read all the other ACK ranges + for range numBlocks { + g, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + gap := protocol.PacketNumber(g) + if smallest < gap+2 { + return 0, errInvalidAckRanges + } + largest := smallest - gap - 2 + + ab, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + ackBlock := protocol.PacketNumber(ab) + + if ackBlock > largest { + return 0, errInvalidAckRanges + } + smallest = largest - ackBlock + frame.AckRanges = append(frame.AckRanges, AckRange{Smallest: smallest, Largest: largest}) + } + + if !frame.validateAckRanges() { + return 0, errInvalidAckRanges + } + + if ecn { + ect0, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + frame.ECT0 = ect0 + ect1, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + frame.ECT1 = ect1 + ecnce, l, err := quicvarint.Parse(b) + if err != nil { + return 0, replaceUnexpectedEOF(err) + } + b = b[l:] + frame.ECNCE = ecnce + } + + return startLen - len(b), nil +} + +// Append appends an ACK frame. +func (f *AckFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + hasECN := f.ECT0 > 0 || f.ECT1 > 0 || f.ECNCE > 0 + if hasECN { + b = append(b, byte(FrameTypeAckECN)) + } else { + b = append(b, byte(FrameTypeAck)) + } + b = quicvarint.Append(b, uint64(f.LargestAcked())) + b = quicvarint.Append(b, encodeAckDelay(f.DelayTime)) + + numRanges := min(len(f.AckRanges), protocol.MaxNumAckRanges) + b = quicvarint.Append(b, uint64(numRanges-1)) + + // write the first range + _, firstRange := f.encodeAckRange(0) + b = quicvarint.Append(b, firstRange) + + // write all the other range + for i := 1; i < numRanges; i++ { + gap, len := f.encodeAckRange(i) + b = quicvarint.Append(b, gap) + b = quicvarint.Append(b, len) + } + + if hasECN { + b = quicvarint.Append(b, f.ECT0) + b = quicvarint.Append(b, f.ECT1) + b = quicvarint.Append(b, f.ECNCE) + } + return b, nil +} + +// Length of a written frame +func (f *AckFrame) Length(_ protocol.Version) protocol.ByteCount { + largestAcked := f.AckRanges[0].Largest + + // The number of ACK ranges is limited to 64, which guarantees that the + // ACK Range Count value can be encoded in a single byte varint. + length := 1 + quicvarint.Len(uint64(largestAcked)) + quicvarint.Len(encodeAckDelay(f.DelayTime)) + 1 + + lowestInFirstRange := f.AckRanges[0].Smallest + length += quicvarint.Len(uint64(largestAcked - lowestInFirstRange)) + + for i := 1; i < min(len(f.AckRanges), protocol.MaxNumAckRanges); i++ { + gap, len := f.encodeAckRange(i) + length += quicvarint.Len(gap) + length += quicvarint.Len(len) + } + if f.ECT0 > 0 || f.ECT1 > 0 || f.ECNCE > 0 { + length += quicvarint.Len(f.ECT0) + quicvarint.Len(f.ECT1) + quicvarint.Len(f.ECNCE) + } + return protocol.ByteCount(length) +} + +// Truncate truncates the ACK frame to fit into maxSize, +// and to at most 64 ACK ranges. +// maxSize must be large enough to fit at least one ACK range. +func (f *AckFrame) Truncate(maxSize protocol.ByteCount, _ protocol.Version) { + f.AckRanges = f.AckRanges[:f.numEncodableAckRanges(maxSize)] +} + +// gets the number of ACK ranges that can be encoded +// such that the resulting frame is smaller than maxSize +func (f *AckFrame) numEncodableAckRanges(maxSize protocol.ByteCount) int { + // Fast path: Most ACK frames are relatively small, and we don't need to calculate the exact length. + // We just assume the worst case scenario: every varint is encoded to 8 bytes. + // If the result is still smaller than the maximum ACK frame size, the actual ACK frame will definitely fit. + length := 1 + 8 /* largest acked */ + 8 /* delay */ + 1 /* ack range count */ + 8 /* first range */ + if f.ECT0 > 0 || f.ECT1 > 0 || f.ECNCE > 0 { + length += 8 + 8 + 8 + } + numRanges := min(len(f.AckRanges), protocol.MaxNumAckRanges) + length += 2 * 8 * (numRanges - 1) + if protocol.ByteCount(length) <= maxSize { + return numRanges + } + + // Slow path: Calculate the exact length of the ACK frame. + length = 1 + quicvarint.Len(uint64(f.LargestAcked())) + quicvarint.Len(encodeAckDelay(f.DelayTime)) + 1 + _, firstRange := f.encodeAckRange(0) + length += quicvarint.Len(firstRange) + if f.ECT0 > 0 || f.ECT1 > 0 || f.ECNCE > 0 { + length += quicvarint.Len(f.ECT0) + quicvarint.Len(f.ECT1) + quicvarint.Len(f.ECNCE) + } + for i := 1; i < numRanges; i++ { + gap, l := f.encodeAckRange(i) + rangeLen := quicvarint.Len(gap) + quicvarint.Len(l) + if protocol.ByteCount(length+rangeLen) > maxSize { + // Writing range i would exceed the maximum size, + // so encode one range less than that. + return i + } + length += rangeLen + } + return numRanges +} + +func (f *AckFrame) encodeAckRange(i int) (gap, length uint64) { + if i == 0 { + return 0, uint64(f.AckRanges[0].Largest - f.AckRanges[0].Smallest) + } + return uint64(f.AckRanges[i-1].Smallest - f.AckRanges[i].Largest - 2), + uint64(f.AckRanges[i].Largest - f.AckRanges[i].Smallest) +} + +// HasMissingRanges returns if this frame reports any missing packets +func (f *AckFrame) HasMissingRanges() bool { + return len(f.AckRanges) > 1 +} + +func (f *AckFrame) validateAckRanges() bool { + if len(f.AckRanges) == 0 { + return false + } + + // check the validity of every single ACK range + for _, ackRange := range f.AckRanges { + if ackRange.Smallest > ackRange.Largest { + return false + } + } + + // check the consistency for ACK with multiple ACK ranges + for i, ackRange := range f.AckRanges { + if i == 0 { + continue + } + lastAckRange := f.AckRanges[i-1] + if lastAckRange.Smallest <= ackRange.Smallest { + return false + } + if lastAckRange.Smallest <= ackRange.Largest+1 { + return false + } + } + + return true +} + +// LargestAcked is the largest acked packet number +func (f *AckFrame) LargestAcked() protocol.PacketNumber { + return f.AckRanges[0].Largest +} + +// LowestAcked is the lowest acked packet number +func (f *AckFrame) LowestAcked() protocol.PacketNumber { + return f.AckRanges[len(f.AckRanges)-1].Smallest +} + +// AcksPacket determines if this ACK frame acks a certain packet number +func (f *AckFrame) AcksPacket(p protocol.PacketNumber) bool { + if p < f.LowestAcked() || p > f.LargestAcked() { + return false + } + + i := sort.Search(len(f.AckRanges), func(i int) bool { + return p >= f.AckRanges[i].Smallest + }) + // i will always be < len(f.AckRanges), since we checked above that p is not bigger than the largest acked + return p <= f.AckRanges[i].Largest +} + +func (f *AckFrame) Reset() { + f.DelayTime = 0 + f.ECT0 = 0 + f.ECT1 = 0 + f.ECNCE = 0 + for _, r := range f.AckRanges { + r.Largest = 0 + r.Smallest = 0 + } + f.AckRanges = f.AckRanges[:0] +} + +func encodeAckDelay(delay time.Duration) uint64 { + return uint64(delay.Nanoseconds() / (1000 * (1 << protocol.AckDelayExponent))) +} diff --git a/third_party/quic-go/internal/wire/ack_frame_test.go b/third_party/quic-go/internal/wire/ack_frame_test.go new file mode 100644 index 0000000..58fe63a --- /dev/null +++ b/third_party/quic-go/internal/wire/ack_frame_test.go @@ -0,0 +1,612 @@ +package wire + +import ( + "io" + "math" + "slices" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + "github.com/stretchr/testify/require" +) + +func TestParseACKWithoutRanges(t *testing.T) { + data := encodeVarInt(100) // largest acked + data = append(data, encodeVarInt(0)...) // delay + data = append(data, encodeVarInt(0)...) // num blocks + data = append(data, encodeVarInt(10)...) // first ack block + var frame AckFrame + n, err := parseAckFrame(&frame, data, FrameTypeAck, protocol.AckDelayExponent, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), n) + require.Equal(t, protocol.PacketNumber(100), frame.LargestAcked()) + require.Equal(t, protocol.PacketNumber(90), frame.LowestAcked()) + require.False(t, frame.HasMissingRanges()) +} + +func TestParseACKSinglePacket(t *testing.T) { + data := encodeVarInt(55) // largest acked + data = append(data, encodeVarInt(0)...) // delay + data = append(data, encodeVarInt(0)...) // num blocks + data = append(data, encodeVarInt(0)...) // first ack block + var frame AckFrame + n, err := parseAckFrame(&frame, data, FrameTypeAck, protocol.AckDelayExponent, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), n) + require.Equal(t, protocol.PacketNumber(55), frame.LargestAcked()) + require.Equal(t, protocol.PacketNumber(55), frame.LowestAcked()) + require.False(t, frame.HasMissingRanges()) +} + +func TestParseACKAllPacketsFrom0ToLargest(t *testing.T) { + data := encodeVarInt(20) // largest acked + data = append(data, encodeVarInt(0)...) // delay + data = append(data, encodeVarInt(0)...) // num blocks + data = append(data, encodeVarInt(20)...) // first ack block + var frame AckFrame + n, err := parseAckFrame(&frame, data, FrameTypeAck, protocol.AckDelayExponent, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), n) + require.Equal(t, protocol.PacketNumber(20), frame.LargestAcked()) + require.Equal(t, protocol.PacketNumber(0), frame.LowestAcked()) + require.False(t, frame.HasMissingRanges()) +} + +func TestParseACKRejectFirstBlockLargerThanLargestAcked(t *testing.T) { + data := encodeVarInt(20) // largest acked + data = append(data, encodeVarInt(0)...) // delay + data = append(data, encodeVarInt(0)...) // num blocks + data = append(data, encodeVarInt(21)...) // first ack block + var frame AckFrame + _, err := parseAckFrame(&frame, data, FrameTypeAck, protocol.AckDelayExponent, protocol.Version1) + require.EqualError(t, err, "invalid first ACK range") +} + +func TestParseACKWithSingleBlock(t *testing.T) { + data := encodeVarInt(1000) // largest acked + data = append(data, encodeVarInt(0)...) // delay + data = append(data, encodeVarInt(1)...) // num blocks + data = append(data, encodeVarInt(100)...) // first ack block + data = append(data, encodeVarInt(98)...) // gap + data = append(data, encodeVarInt(50)...) // ack block + var frame AckFrame + n, err := parseAckFrame(&frame, data, FrameTypeAck, protocol.AckDelayExponent, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), n) + require.Equal(t, protocol.PacketNumber(1000), frame.LargestAcked()) + require.Equal(t, protocol.PacketNumber(750), frame.LowestAcked()) + require.True(t, frame.HasMissingRanges()) + require.Equal(t, []AckRange{ + {Largest: 1000, Smallest: 900}, + {Largest: 800, Smallest: 750}, + }, frame.AckRanges) +} + +func TestParseACKWithMultipleBlocks(t *testing.T) { + data := encodeVarInt(100) // largest acked + data = append(data, encodeVarInt(0)...) // delay + data = append(data, encodeVarInt(2)...) // num blocks + data = append(data, encodeVarInt(0)...) // first ack block + data = append(data, encodeVarInt(0)...) // gap + data = append(data, encodeVarInt(0)...) // ack block + data = append(data, encodeVarInt(1)...) // gap + data = append(data, encodeVarInt(1)...) // ack block + var frame AckFrame + n, err := parseAckFrame(&frame, data, FrameTypeAck, protocol.AckDelayExponent, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), n) + require.Equal(t, protocol.PacketNumber(100), frame.LargestAcked()) + require.Equal(t, protocol.PacketNumber(94), frame.LowestAcked()) + require.True(t, frame.HasMissingRanges()) + require.Equal(t, []AckRange{ + {Largest: 100, Smallest: 100}, + {Largest: 98, Smallest: 98}, + {Largest: 95, Smallest: 94}, + }, frame.AckRanges) +} + +func TestParseACKUseAckDelayExponent(t *testing.T) { + const delayTime = 1 << 10 * time.Millisecond + f := &AckFrame{ + AckRanges: []AckRange{{Smallest: 1, Largest: 1}}, + DelayTime: delayTime, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + for i := range uint8(8) { + typ, l, err := quicvarint.Parse(b) + require.NoError(t, err) + var frame AckFrame + n, err := parseAckFrame(&frame, b[l:], FrameType(typ), protocol.AckDelayExponent+i, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(b[l:]), n) + require.Equal(t, delayTime*(1< len(b) { + return nil, 0, io.EOF + } + + reasonPhrase := make([]byte, reasonPhraseLen) + copy(reasonPhrase, b) + f.ReasonPhrase = string(reasonPhrase) + return f, startLen - len(b) + int(reasonPhraseLen), nil +} + +// Length of a written frame +func (f *ConnectionCloseFrame) Length(protocol.Version) protocol.ByteCount { + length := 1 + protocol.ByteCount(quicvarint.Len(f.ErrorCode)+quicvarint.Len(uint64(len(f.ReasonPhrase)))) + protocol.ByteCount(len(f.ReasonPhrase)) + if !f.IsApplicationError { + length += protocol.ByteCount(quicvarint.Len(f.FrameType)) // for the frame type + } + return length +} + +func (f *ConnectionCloseFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + if f.IsApplicationError { + b = append(b, byte(FrameTypeApplicationClose)) + } else { + b = append(b, byte(FrameTypeConnectionClose)) + } + + b = quicvarint.Append(b, f.ErrorCode) + if !f.IsApplicationError { + b = quicvarint.Append(b, f.FrameType) + } + b = quicvarint.Append(b, uint64(len(f.ReasonPhrase))) + b = append(b, []byte(f.ReasonPhrase)...) + return b, nil +} diff --git a/third_party/quic-go/internal/wire/connection_close_frame_test.go b/third_party/quic-go/internal/wire/connection_close_frame_test.go new file mode 100644 index 0000000..2831e9b --- /dev/null +++ b/third_party/quic-go/internal/wire/connection_close_frame_test.go @@ -0,0 +1,137 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseConnectionCloseTransportError(t *testing.T) { + reason := "No recent network activity." + data := encodeVarInt(0x19) + data = append(data, encodeVarInt(0x1337)...) // frame type + data = append(data, encodeVarInt(uint64(len(reason)))...) // reason phrase length + data = append(data, []byte(reason)...) + frame, l, err := parseConnectionCloseFrame(data, FrameTypeConnectionClose, protocol.Version1) + require.NoError(t, err) + require.False(t, frame.IsApplicationError) + require.EqualValues(t, 0x19, frame.ErrorCode) + require.EqualValues(t, 0x1337, frame.FrameType) + require.Equal(t, reason, frame.ReasonPhrase) + require.Equal(t, len(data), l) +} + +func TestParseConnectionCloseWithApplicationError(t *testing.T) { + reason := "The application messed things up." + data := encodeVarInt(0xcafe) + data = append(data, encodeVarInt(uint64(len(reason)))...) // reason phrase length + data = append(data, reason...) + frame, l, err := parseConnectionCloseFrame(data, FrameTypeApplicationClose, protocol.Version1) + require.NoError(t, err) + require.True(t, frame.IsApplicationError) + require.EqualValues(t, 0xcafe, frame.ErrorCode) + require.Equal(t, reason, frame.ReasonPhrase) + require.Equal(t, len(data), l) +} + +func TestParseConnectionCloseLongReasonPhrase(t *testing.T) { + data := encodeVarInt(0xcafe) + data = append(data, encodeVarInt(0x42)...) // frame type + data = append(data, encodeVarInt(0xffff)...) // reason phrase length + _, _, err := parseConnectionCloseFrame(data, FrameTypeConnectionClose, protocol.Version1) + require.Equal(t, io.EOF, err) +} + +func TestParseConnectionCloseErrorsOnEOFs(t *testing.T) { + reason := "No recent network activity." + data := encodeVarInt(0x19) + data = append(data, encodeVarInt(0x1337)...) // frame type + data = append(data, encodeVarInt(uint64(len(reason)))...) // reason phrase length + data = append(data, []byte(reason)...) + _, l, err := parseConnectionCloseFrame(data, FrameTypeConnectionClose, protocol.Version1) + require.Equal(t, len(data), l) + require.NoError(t, err) + for i := range data { + _, _, err = parseConnectionCloseFrame(data[:i], FrameTypeConnectionClose, protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestParseConnectionCloseNoReasonPhrase(t *testing.T) { + data := encodeVarInt(0xcafe) + data = append(data, encodeVarInt(0x42)...) // frame type + data = append(data, encodeVarInt(0)...) + frame, l, err := parseConnectionCloseFrame(data, FrameTypeConnectionClose, protocol.Version1) + require.NoError(t, err) + require.Empty(t, frame.ReasonPhrase) + require.Equal(t, len(data), l) +} + +func TestWriteConnectionCloseNoReasonPhrase(t *testing.T) { + frame := &ConnectionCloseFrame{ + ErrorCode: 0xbeef, + FrameType: 0x12345, + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeConnectionClose)} + expected = append(expected, encodeVarInt(0xbeef)...) + expected = append(expected, encodeVarInt(0x12345)...) // frame type + expected = append(expected, encodeVarInt(0)...) // reason phrase length + require.Equal(t, expected, b) +} + +func TestWriteConnectionCloseWithReasonPhrase(t *testing.T) { + frame := &ConnectionCloseFrame{ + ErrorCode: 0xdead, + ReasonPhrase: "foobar", + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeConnectionClose)} + expected = append(expected, encodeVarInt(0xdead)...) + expected = append(expected, encodeVarInt(0)...) // frame type + expected = append(expected, encodeVarInt(6)...) // reason phrase length + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) +} + +func TestWriteConnectionCloseWithApplicationError(t *testing.T) { + frame := &ConnectionCloseFrame{ + IsApplicationError: true, + ErrorCode: 0xdead, + ReasonPhrase: "foobar", + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeApplicationClose)} + expected = append(expected, encodeVarInt(0xdead)...) + expected = append(expected, encodeVarInt(6)...) // reason phrase length + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) +} + +func TestWriteConnectionCloseTransportError(t *testing.T) { + f := &ConnectionCloseFrame{ + ErrorCode: 0xcafe, + FrameType: 0xdeadbeef, + ReasonPhrase: "foobar", + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Len(t, b, int(f.Length(protocol.Version1))) +} + +func TestWriteConnectionCloseLength(t *testing.T) { + f := &ConnectionCloseFrame{ + IsApplicationError: true, + ErrorCode: 0xcafe, + ReasonPhrase: "foobar", + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Len(t, b, int(f.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/crypto_frame.go b/third_party/quic-go/internal/wire/crypto_frame.go new file mode 100644 index 0000000..b792e54 --- /dev/null +++ b/third_party/quic-go/internal/wire/crypto_frame.go @@ -0,0 +1,97 @@ +package wire + +import ( + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A CryptoFrame is a CRYPTO frame +type CryptoFrame struct { + Offset protocol.ByteCount + Data []byte +} + +func parseCryptoFrame(b []byte, _ protocol.Version) (*CryptoFrame, int, error) { + startLen := len(b) + frame := &CryptoFrame{} + offset, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + frame.Offset = protocol.ByteCount(offset) + dataLen, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + if dataLen > uint64(len(b)) { + return nil, 0, io.EOF + } + if dataLen != 0 { + frame.Data = make([]byte, dataLen) + copy(frame.Data, b) + } + return frame, startLen - len(b) + int(dataLen), nil +} + +func (f *CryptoFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeCrypto)) + b = quicvarint.Append(b, uint64(f.Offset)) + b = quicvarint.Append(b, uint64(len(f.Data))) + b = append(b, f.Data...) + return b, nil +} + +// Length of a written frame +func (f *CryptoFrame) Length(_ protocol.Version) protocol.ByteCount { + return protocol.ByteCount(1 + quicvarint.Len(uint64(f.Offset)) + quicvarint.Len(uint64(len(f.Data))) + len(f.Data)) +} + +// MaxDataLen returns the maximum data length +func (f *CryptoFrame) MaxDataLen(maxSize protocol.ByteCount) protocol.ByteCount { + // pretend that the data size will be 1 bytes + // if it turns out that varint encoding the length will consume 2 bytes, we need to adjust the data length afterwards + headerLen := protocol.ByteCount(1 + quicvarint.Len(uint64(f.Offset)) + 1) + if headerLen > maxSize { + return 0 + } + maxDataLen := maxSize - headerLen + if quicvarint.Len(uint64(maxDataLen)) != 1 { + maxDataLen-- + } + return maxDataLen +} + +// MaybeSplitOffFrame splits a frame such that it is not bigger than n bytes. +// It returns if the frame was actually split. +// The frame might not be split if: +// * the size is large enough to fit the whole frame +// * the size is too small to fit even a 1-byte frame. In that case, the frame returned is nil. +func (f *CryptoFrame) MaybeSplitOffFrame(maxSize protocol.ByteCount, version protocol.Version) (*CryptoFrame, bool /* was splitting required */) { + if f.Length(version) <= maxSize { + return nil, false + } + + n := f.MaxDataLen(maxSize) + if n == 0 { + return nil, true + } + + newLen := protocol.ByteCount(len(f.Data)) - n + + new := &CryptoFrame{} + new.Offset = f.Offset + new.Data = make([]byte, newLen) + + // swap the data slices + new.Data, f.Data = f.Data, new.Data + + copy(f.Data, new.Data[n:]) + new.Data = new.Data[:n] + f.Offset += n + + return new, true +} diff --git a/third_party/quic-go/internal/wire/crypto_frame_test.go b/third_party/quic-go/internal/wire/crypto_frame_test.go new file mode 100644 index 0000000..76608d8 --- /dev/null +++ b/third_party/quic-go/internal/wire/crypto_frame_test.go @@ -0,0 +1,123 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseCryptoFrame(t *testing.T) { + data := encodeVarInt(0xdecafbad) // offset + data = append(data, encodeVarInt(6)...) // length + data = append(data, []byte("foobar")...) + frame, l, err := parseCryptoFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.ByteCount(0xdecafbad), frame.Offset) + require.Equal(t, []byte("foobar"), frame.Data) + require.Equal(t, len(data), l) +} + +func TestParseCryptoFrameErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdecafbad) // offset + data = append(data, encodeVarInt(6)...) // data length + data = append(data, []byte("foobar")...) + _, l, err := parseCryptoFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseCryptoFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteCryptoFrame(t *testing.T) { + f := &CryptoFrame{ + Offset: 0x123456, + Data: []byte("foobar"), + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeCrypto)} + expected = append(expected, encodeVarInt(0x123456)...) // offset + expected = append(expected, encodeVarInt(6)...) // length + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestCryptoFrameMaxDataLength(t *testing.T) { + const maxSize = 3000 + + data := make([]byte, maxSize) + f := &CryptoFrame{ + Offset: 0xdeadbeef, + } + var frameOneByteTooSmallCounter int + for i := 1; i < maxSize; i++ { + f.Data = nil + maxDataLen := f.MaxDataLen(protocol.ByteCount(i)) + if maxDataLen == 0 { // 0 means that no valid CRYPTO frame can be written + // check that writing a minimal size CRYPTO frame (i.e. with 1 byte data) is actually larger than the desired size + f.Data = []byte{0} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Greater(t, len(b), i) + continue + } + f.Data = data[:int(maxDataLen)] + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + // There's *one* pathological case, where a data length of x can be encoded into 1 byte + // but a data lengths of x+1 needs 2 bytes + // In that case, it's impossible to create a STREAM frame of the desired size + if len(b) == i-1 { + frameOneByteTooSmallCounter++ + continue + } + require.Equal(t, i, len(b)) + } + require.Equal(t, 1, frameOneByteTooSmallCounter) +} + +func TestCryptoFrameSplitting(t *testing.T) { + f := &CryptoFrame{ + Offset: 0x1337, + Data: []byte("foobar"), + } + hdrLen := f.Length(protocol.Version1) - 6 + new, needsSplit := f.MaybeSplitOffFrame(hdrLen+3, protocol.Version1) + require.True(t, needsSplit) + require.Equal(t, []byte("foo"), new.Data) + require.Equal(t, protocol.ByteCount(0x1337), new.Offset) + require.Equal(t, []byte("bar"), f.Data) + require.Equal(t, protocol.ByteCount(0x1337+3), f.Offset) +} + +func TestCryptoFrameNoSplitWhenEnoughSpace(t *testing.T) { + f := &CryptoFrame{ + Offset: 0x1337, + Data: []byte("foobar"), + } + splitFrame, needsSplit := f.MaybeSplitOffFrame(f.Length(protocol.Version1), protocol.Version1) + require.False(t, needsSplit) + require.Nil(t, splitFrame) +} + +func TestCryptoFrameNoSplitWhenSizeTooSmall(t *testing.T) { + f := &CryptoFrame{ + Offset: 0x1337, + Data: []byte("foobar"), + } + length := f.Length(protocol.Version1) - 6 + for i := protocol.ByteCount(0); i <= length; i++ { + splitFrame, needsSplit := f.MaybeSplitOffFrame(i, protocol.Version1) + require.True(t, needsSplit) + require.Nil(t, splitFrame) + } + splitFrame, needsSplit := f.MaybeSplitOffFrame(length+1, protocol.Version1) + require.True(t, needsSplit) + require.NotNil(t, splitFrame) +} diff --git a/third_party/quic-go/internal/wire/data_blocked_frame.go b/third_party/quic-go/internal/wire/data_blocked_frame.go new file mode 100644 index 0000000..1dc51db --- /dev/null +++ b/third_party/quic-go/internal/wire/data_blocked_frame.go @@ -0,0 +1,29 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A DataBlockedFrame is a DATA_BLOCKED frame +type DataBlockedFrame struct { + MaximumData protocol.ByteCount +} + +func parseDataBlockedFrame(b []byte, _ protocol.Version) (*DataBlockedFrame, int, error) { + offset, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + return &DataBlockedFrame{MaximumData: protocol.ByteCount(offset)}, l, nil +} + +func (f *DataBlockedFrame) Append(b []byte, version protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeDataBlocked)) + return quicvarint.Append(b, uint64(f.MaximumData)), nil +} + +// Length of a written frame +func (f *DataBlockedFrame) Length(version protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.MaximumData))) +} diff --git a/third_party/quic-go/internal/wire/data_blocked_frame_test.go b/third_party/quic-go/internal/wire/data_blocked_frame_test.go new file mode 100644 index 0000000..d7d742f --- /dev/null +++ b/third_party/quic-go/internal/wire/data_blocked_frame_test.go @@ -0,0 +1,40 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +func TestParseDataBlocked(t *testing.T) { + data := encodeVarInt(0x12345678) + frame, l, err := parseDataBlockedFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.ByteCount(0x12345678), frame.MaximumData) + require.Equal(t, len(data), l) +} + +func TestParseDataBlockedErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0x12345678) + _, l, err := parseDataBlockedFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseDataBlockedFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteDataBlocked(t *testing.T) { + frame := DataBlockedFrame{MaximumData: 0xdeadbeef} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeDataBlocked)} + expected = append(expected, encodeVarInt(0xdeadbeef)...) + require.Equal(t, expected, b) + require.Equal(t, protocol.ByteCount(1+quicvarint.Len(uint64(frame.MaximumData))), frame.Length(protocol.Version1)) +} diff --git a/third_party/quic-go/internal/wire/datagram_frame.go b/third_party/quic-go/internal/wire/datagram_frame.go new file mode 100644 index 0000000..a9fcdcb --- /dev/null +++ b/third_party/quic-go/internal/wire/datagram_frame.go @@ -0,0 +1,85 @@ +package wire + +import ( + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// MaxDatagramSize is the maximum size of a DATAGRAM frame (RFC 9221). +// By setting it to a large value, we allow all datagrams that fit into a QUIC packet. +// The value is chosen such that it can still be encoded as a 2 byte varint. +// This is a var and not a const so it can be set in tests. +var MaxDatagramSize protocol.ByteCount = 16383 + +// A DatagramFrame is a DATAGRAM frame +type DatagramFrame struct { + DataLenPresent bool + Data []byte +} + +func parseDatagramFrame(b []byte, typ FrameType, _ protocol.Version) (*DatagramFrame, int, error) { + startLen := len(b) + f := &DatagramFrame{} + f.DataLenPresent = uint64(typ)&0x1 > 0 + + var length uint64 + if f.DataLenPresent { + var err error + var l int + length, l, err = quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + if length > uint64(len(b)) { + return nil, 0, io.EOF + } + } else { + length = uint64(len(b)) + } + f.Data = make([]byte, length) + copy(f.Data, b) + return f, startLen - len(b) + int(length), nil +} + +func (f *DatagramFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + typ := uint8(0x30) + if f.DataLenPresent { + typ ^= 0b1 + } + b = append(b, typ) + if f.DataLenPresent { + b = quicvarint.Append(b, uint64(len(f.Data))) + } + b = append(b, f.Data...) + return b, nil +} + +// MaxDataLen returns the maximum data length +func (f *DatagramFrame) MaxDataLen(maxSize protocol.ByteCount, version protocol.Version) protocol.ByteCount { + headerLen := protocol.ByteCount(1) + if f.DataLenPresent { + // pretend that the data size will be 1 bytes + // if it turns out that varint encoding the length will consume 2 bytes, we need to adjust the data length afterwards + headerLen++ + } + if headerLen > maxSize { + return 0 + } + maxDataLen := maxSize - headerLen + if f.DataLenPresent && quicvarint.Len(uint64(maxDataLen)) != 1 { + maxDataLen-- + } + return maxDataLen +} + +// Length of a written frame +func (f *DatagramFrame) Length(_ protocol.Version) protocol.ByteCount { + length := 1 + protocol.ByteCount(len(f.Data)) + if f.DataLenPresent { + length += protocol.ByteCount(quicvarint.Len(uint64(len(f.Data)))) + } + return length +} diff --git a/third_party/quic-go/internal/wire/datagram_frame_test.go b/third_party/quic-go/internal/wire/datagram_frame_test.go new file mode 100644 index 0000000..db04979 --- /dev/null +++ b/third_party/quic-go/internal/wire/datagram_frame_test.go @@ -0,0 +1,126 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseDatagramFrameWithLength(t *testing.T) { + data := encodeVarInt(0x6) // length + data = append(data, []byte("foobar")...) + frame, l, err := parseDatagramFrame(data, 0x30^0x1, protocol.Version1) + require.NoError(t, err) + require.Equal(t, []byte("foobar"), frame.Data) + require.True(t, frame.DataLenPresent) + require.Equal(t, len(data), l) +} + +func TestParseDatagramFrameWithoutLength(t *testing.T) { + data := []byte("Lorem ipsum dolor sit amet") + frame, l, err := parseDatagramFrame(data, 0x30, protocol.Version1) + require.NoError(t, err) + require.Equal(t, []byte("Lorem ipsum dolor sit amet"), frame.Data) + require.False(t, frame.DataLenPresent) + require.Equal(t, len(data), l) +} + +func TestParseDatagramFrameErrorsOnLengthLongerThanFrame(t *testing.T) { + data := encodeVarInt(0x6) // length + data = append(data, []byte("fooba")...) + _, _, err := parseDatagramFrame(data, 0x30^0x1, protocol.Version1) + require.Equal(t, io.EOF, err) +} + +func TestParseDatagramFrameErrorsOnEOFs(t *testing.T) { + const typ = 0x30 ^ 0x1 + data := encodeVarInt(6) // length + data = append(data, []byte("foobar")...) + _, l, err := parseDatagramFrame(data, typ, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err = parseDatagramFrame(data[0:i], typ, protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteDatagramFrameWithLength(t *testing.T) { + f := &DatagramFrame{ + DataLenPresent: true, + Data: []byte("foobar"), + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x30 ^ 0x1} + expected = append(expected, encodeVarInt(0x6)...) + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) + require.Len(t, b, int(f.Length(protocol.Version1))) +} + +func TestWriteDatagramFrameWithoutLength(t *testing.T) { + f := &DatagramFrame{Data: []byte("Lorem ipsum")} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x30} + expected = append(expected, []byte("Lorem ipsum")...) + require.Equal(t, expected, b) + require.Len(t, b, int(f.Length(protocol.Version1))) +} + +func TestMaxDatagramLenWithoutDataLenPresent(t *testing.T) { + const maxSize = 3000 + data := make([]byte, maxSize) + f := &DatagramFrame{} + for i := 1; i < 3000; i++ { + f.Data = nil + maxDataLen := f.MaxDataLen(protocol.ByteCount(i), protocol.Version1) + if maxDataLen == 0 { // 0 means that no valid DATAGRAM frame can be written + // check that writing a minimal size DATAGRAM frame (i.e. with 1 byte data) is actually larger than the desired size + f.Data = []byte{0} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Greater(t, len(b), i) + continue + } + f.Data = data[:int(maxDataLen)] + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Len(t, b, i) + } +} + +func TestMaxDatagramLenWithDataLenPresent(t *testing.T) { + const maxSize = 3000 + data := make([]byte, maxSize) + f := &DatagramFrame{DataLenPresent: true} + var frameOneByteTooSmallCounter int + for i := 1; i < 3000; i++ { + f.Data = nil + maxDataLen := f.MaxDataLen(protocol.ByteCount(i), protocol.Version1) + if maxDataLen == 0 { // 0 means that no valid DATAGRAM frame can be written + // check that writing a minimal size DATAGRAM frame (i.e. with 1 byte data) is actually larger than the desired size + f.Data = []byte{0} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Greater(t, len(b), i) + continue + } + f.Data = data[:int(maxDataLen)] + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + // There's *one* pathological case, where a data length of x can be encoded into 1 byte + // but a data lengths of x+1 needs 2 bytes + // In that case, it's impossible to create a DATAGRAM frame of the desired size + if len(b) == i-1 { + frameOneByteTooSmallCounter++ + continue + } + require.Len(t, b, i) + } + require.Equal(t, 1, frameOneByteTooSmallCounter) +} diff --git a/third_party/quic-go/internal/wire/extended_header.go b/third_party/quic-go/internal/wire/extended_header.go new file mode 100644 index 0000000..1a8c439 --- /dev/null +++ b/third_party/quic-go/internal/wire/extended_header.go @@ -0,0 +1,164 @@ +package wire + +import ( + "encoding/binary" + "errors" + "fmt" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/quicvarint" +) + +// ErrInvalidReservedBits is returned when the reserved bits are incorrect. +// When this error is returned, parsing continues, and an ExtendedHeader is returned. +// This is necessary because we need to decrypt the packet in that case, +// in order to avoid a timing side-channel. +var ErrInvalidReservedBits = errors.New("invalid reserved bits") + +// ExtendedHeader is the header of a QUIC packet. +type ExtendedHeader struct { + Header + + typeByte byte + + KeyPhase protocol.KeyPhaseBit + + PacketNumberLen protocol.PacketNumberLen + PacketNumber protocol.PacketNumber + + parsedLen protocol.ByteCount +} + +func (h *ExtendedHeader) parse(data []byte) (bool /* reserved bits valid */, error) { + // read the (now unencrypted) first byte + h.typeByte = data[0] + h.PacketNumberLen = protocol.PacketNumberLen(h.typeByte&0x3) + 1 + if protocol.ByteCount(len(data)) < h.Header.ParsedLen()+protocol.ByteCount(h.PacketNumberLen) { + return false, io.EOF + } + + pn, err := readPacketNumber(data[h.Header.ParsedLen():], h.PacketNumberLen) + if err != nil { + return true, nil + } + h.PacketNumber = pn + reservedBitsValid := h.typeByte&0xc == 0 + + h.parsedLen = h.Header.ParsedLen() + protocol.ByteCount(h.PacketNumberLen) + return reservedBitsValid, err +} + +// Append appends the Header. +func (h *ExtendedHeader) Append(b []byte, v protocol.Version) ([]byte, error) { + if h.DestConnectionID.Len() > protocol.MaxConnIDLen { + return nil, fmt.Errorf("invalid connection ID length: %d bytes", h.DestConnectionID.Len()) + } + if h.SrcConnectionID.Len() > protocol.MaxConnIDLen { + return nil, fmt.Errorf("invalid connection ID length: %d bytes", h.SrcConnectionID.Len()) + } + + var packetType uint8 + if v == protocol.Version2 { + switch h.Type { + case protocol.PacketTypeInitial: + packetType = 0b01 + case protocol.PacketType0RTT: + packetType = 0b10 + case protocol.PacketTypeHandshake: + packetType = 0b11 + case protocol.PacketTypeRetry: + packetType = 0b00 + } + } else { + switch h.Type { + case protocol.PacketTypeInitial: + packetType = 0b00 + case protocol.PacketType0RTT: + packetType = 0b01 + case protocol.PacketTypeHandshake: + packetType = 0b10 + case protocol.PacketTypeRetry: + packetType = 0b11 + } + } + firstByte := 0xc0 | packetType<<4 + if h.Type != protocol.PacketTypeRetry { + // Retry packets don't have a packet number + firstByte |= uint8(h.PacketNumberLen - 1) + } + + b = append(b, firstByte) + b = append(b, make([]byte, 4)...) + binary.BigEndian.PutUint32(b[len(b)-4:], uint32(h.Version)) + b = append(b, uint8(h.DestConnectionID.Len())) + b = append(b, h.DestConnectionID.Bytes()...) + b = append(b, uint8(h.SrcConnectionID.Len())) + b = append(b, h.SrcConnectionID.Bytes()...) + + //nolint:exhaustive + switch h.Type { + case protocol.PacketTypeRetry: + b = append(b, h.Token...) + return b, nil + case protocol.PacketTypeInitial: + b = quicvarint.Append(b, uint64(len(h.Token))) + b = append(b, h.Token...) + } + b = quicvarint.AppendWithLen(b, uint64(h.Length), 2) + return appendPacketNumber(b, h.PacketNumber, h.PacketNumberLen) +} + +// ParsedLen returns the number of bytes that were consumed when parsing the header +func (h *ExtendedHeader) ParsedLen() protocol.ByteCount { + return h.parsedLen +} + +// GetLength determines the length of the Header. +func (h *ExtendedHeader) GetLength(_ protocol.Version) protocol.ByteCount { + length := 1 /* type byte */ + 4 /* version */ + 1 /* dest conn ID len */ + protocol.ByteCount(h.DestConnectionID.Len()) + 1 /* src conn ID len */ + protocol.ByteCount(h.SrcConnectionID.Len()) + protocol.ByteCount(h.PacketNumberLen) + 2 /* length */ + if h.Type == protocol.PacketTypeInitial { + length += protocol.ByteCount(quicvarint.Len(uint64(len(h.Token))) + len(h.Token)) + } + return length +} + +// Log logs the Header +func (h *ExtendedHeader) Log(logger utils.Logger) { + var token string + if h.Type == protocol.PacketTypeInitial || h.Type == protocol.PacketTypeRetry { + if len(h.Token) == 0 { + token = "Token: (empty), " + } else { + token = fmt.Sprintf("Token: %#x, ", h.Token) + } + if h.Type == protocol.PacketTypeRetry { + logger.Debugf("\tLong Header{Type: %s, DestConnectionID: %s, SrcConnectionID: %s, %sVersion: %s}", h.Type, h.DestConnectionID, h.SrcConnectionID, token, h.Version) + return + } + } + logger.Debugf("\tLong Header{Type: %s, DestConnectionID: %s, SrcConnectionID: %s, %sPacketNumber: %d, PacketNumberLen: %d, Length: %d, Version: %s}", h.Type, h.DestConnectionID, h.SrcConnectionID, token, h.PacketNumber, h.PacketNumberLen, h.Length, h.Version) +} + +func appendPacketNumber(b []byte, pn protocol.PacketNumber, pnLen protocol.PacketNumberLen) ([]byte, error) { + switch pnLen { + case protocol.PacketNumberLen1: + b = append(b, uint8(pn)) + case protocol.PacketNumberLen2: + buf := make([]byte, 2) + binary.BigEndian.PutUint16(buf, uint16(pn)) + b = append(b, buf...) + case protocol.PacketNumberLen3: + buf := make([]byte, 4) + binary.BigEndian.PutUint32(buf, uint32(pn)) + b = append(b, buf[1:]...) + case protocol.PacketNumberLen4: + buf := make([]byte, 4) + binary.BigEndian.PutUint32(buf, uint32(pn)) + b = append(b, buf...) + default: + return nil, fmt.Errorf("invalid packet number length: %d", pnLen) + } + return b, nil +} diff --git a/third_party/quic-go/internal/wire/extended_header_test.go b/third_party/quic-go/internal/wire/extended_header_test.go new file mode 100644 index 0000000..b95d304 --- /dev/null +++ b/third_party/quic-go/internal/wire/extended_header_test.go @@ -0,0 +1,269 @@ +package wire + +import ( + "bytes" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +func TestWritesLongHeaderVersion1(t *testing.T) { + header := &ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad, 0x0, 0x0, 0x13, 0x37}), + Version: 0x1020304, + Length: 1234, + }, + PacketNumber: 0xdecaf, + PacketNumberLen: protocol.PacketNumberLen3, + } + b, err := header.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{ + 0xc0 | 0x2<<4 | 0x2, + 0x1, 0x2, 0x3, 0x4, // version number + 0x6, // dest connection ID length + 0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, // dest connection ID + 0x8, // src connection ID length + 0xde, 0xca, 0xfb, 0xad, 0x0, 0x0, 0x13, 0x37, // source connection ID + } + expected = append(expected, encodeVarInt(1234)...) // length + expected = append(expected, []byte{0xd, 0xec, 0xaf}...) // packet number + require.Equal(t, expected, b) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version1)) +} + +func TestWritesHandshakePacketVersion2(t *testing.T) { + header := &ExtendedHeader{ + Header: Header{ + Version: protocol.Version2, + Type: protocol.PacketTypeHandshake, + }, + PacketNumber: 0xdecafbad, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := header.Append(nil, protocol.Version2) + require.NoError(t, err) + require.Equal(t, byte(0b11), b[0]>>4&0b11) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version2)) +} + +func TestWritesHeaderWith20ByteConnectionID(t *testing.T) { + srcConnID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}) + header := &ExtendedHeader{ + Header: Header{ + SrcConnectionID: srcConnID, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20}), // connection IDs must be at most 20 bytes long + Version: 0x1020304, + Type: 0x5, + }, + PacketNumber: 0xdecafbad, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := header.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Contains(t, string(b), string([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20})) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version1)) +} + +func TestWritesInitialContainingToken(t *testing.T) { + token := []byte("Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.") + header := &ExtendedHeader{ + Header: Header{ + Version: 0x1020304, + Type: protocol.PacketTypeInitial, + Token: token, + }, + PacketNumber: 0xdecafbad, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := header.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, byte(0), b[0]>>4&0b11) + expectedSubstring := append(encodeVarInt(uint64(len(token))), token...) + require.Contains(t, string(b), string(expectedSubstring)) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version1)) +} + +func TestUses2ByteEncodingForLengthOnInitialPackets(t *testing.T) { + header := &ExtendedHeader{ + Header: Header{ + Version: 0x1020304, + Type: protocol.PacketTypeInitial, + Length: 37, + }, + PacketNumber: 0xdecafbad, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := header.Append(nil, protocol.Version1) + require.NoError(t, err) + lengthEncoded := quicvarint.AppendWithLen(nil, 37, 2) + require.Equal(t, lengthEncoded, b[len(b)-6:len(b)-4]) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version1)) +} + +func TestWritesInitialPacketVersion2(t *testing.T) { + header := &ExtendedHeader{ + Header: Header{ + Version: protocol.Version2, + Type: protocol.PacketTypeInitial, + }, + PacketNumber: 0xdecafbad, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := header.Append(nil, protocol.Version2) + require.NoError(t, err) + require.Equal(t, byte(0b01), b[0]>>4&0b11) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version2)) +} + +func TestWrites0RTTPacketVersion2(t *testing.T) { + header := &ExtendedHeader{ + Header: Header{ + Version: protocol.Version2, + Type: protocol.PacketType0RTT, + }, + PacketNumber: 0xdecafbad, + PacketNumberLen: protocol.PacketNumberLen4, + } + b, err := header.Append(nil, protocol.Version2) + require.NoError(t, err) + require.Equal(t, byte(0b10), b[0]>>4&0b11) + require.Equal(t, protocol.ByteCount(len(b)), header.GetLength(protocol.Version2)) +} + +func TestWritesRetryPacket(t *testing.T) { + token := []byte("Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat.") + + for _, version := range []protocol.Version{protocol.Version1, protocol.Version2} { + t.Run(version.String(), func(t *testing.T) { + header := &ExtendedHeader{Header: Header{ + Version: version, + Type: protocol.PacketTypeRetry, + Token: token, + }} + b, err := header.Append(nil, version) + require.NoError(t, err) + + var expected []byte + switch version { + case protocol.Version1: + expected = append(expected, 0xc0|0b11<<4) + case protocol.Version2: + expected = append(expected, 0xc0) + } + + expected = appendVersion(expected, version) + expected = append(expected, 0x0) // dest connection ID length + expected = append(expected, 0x0) // src connection ID length + expected = append(expected, token...) + require.Equal(t, expected, b) + }) + } +} + +func TestLogsLongHeaders(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + + (&ExtendedHeader{ + Header: Header{ + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad, 0x013, 0x37, 0x13, 0x37}), + Type: protocol.PacketTypeHandshake, + Length: 54321, + Version: 0xfeed, + }, + PacketNumber: 1337, + PacketNumberLen: protocol.PacketNumberLen2, + }).Log(logger) + require.Contains(t, buf.String(), "Long Header{Type: Handshake, DestConnectionID: deadbeefcafe1337, SrcConnectionID: decafbad13371337, PacketNumber: 1337, PacketNumberLen: 2, Length: 54321, Version: 0xfeed}") +} + +func TestLogsInitialPacketsWithToken(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + + (&ExtendedHeader{ + Header: Header{ + DestConnectionID: protocol.ParseConnectionID([]byte{0xca, 0xfe, 0x13, 0x37}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + Type: protocol.PacketTypeInitial, + Token: []byte{0xde, 0xad, 0xbe, 0xef}, + Length: 100, + Version: 0xfeed, + }, + PacketNumber: 42, + PacketNumberLen: protocol.PacketNumberLen2, + }).Log(logger) + require.Contains(t, buf.String(), "Long Header{Type: Initial, DestConnectionID: cafe1337, SrcConnectionID: decafbad, Token: 0xdeadbeef, PacketNumber: 42, PacketNumberLen: 2, Length: 100, Version: 0xfeed}") +} + +func TestLogsInitialPacketsWithoutToken(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + + (&ExtendedHeader{ + Header: Header{ + DestConnectionID: protocol.ParseConnectionID([]byte{0xca, 0xfe, 0x13, 0x37}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + Type: protocol.PacketTypeInitial, + Length: 100, + Version: 0xfeed, + }, + PacketNumber: 42, + PacketNumberLen: protocol.PacketNumberLen2, + }).Log(logger) + require.Contains(t, buf.String(), "Long Header{Type: Initial, DestConnectionID: cafe1337, SrcConnectionID: decafbad, Token: (empty), PacketNumber: 42, PacketNumberLen: 2, Length: 100, Version: 0xfeed}") +} + +func TestLogsRetryPacketsWithToken(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + + (&ExtendedHeader{ + Header: Header{ + DestConnectionID: protocol.ParseConnectionID([]byte{0xca, 0xfe, 0x13, 0x37}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + Type: protocol.PacketTypeRetry, + Token: []byte{0x12, 0x34, 0x56}, + Version: 0xfeed, + }, + }).Log(logger) + require.Contains(t, buf.String(), "Long Header{Type: Retry, DestConnectionID: cafe1337, SrcConnectionID: decafbad, Token: 0x123456, Version: 0xfeed}") +} + +func BenchmarkParseExtendedHeader(b *testing.B) { + b.ReportAllocs() + + data, err := (&ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe}), + SrcConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad, 0x0, 0x0, 0x13, 0x37}), + Version: protocol.Version1, + Length: 1234, + }, + PacketNumber: 0xdecaf, + PacketNumberLen: protocol.PacketNumberLen3, + }).Append(nil, protocol.Version1) + if err != nil { + b.Fatal(err) + } + data = append(data, make([]byte, 1231)...) + + for b.Loop() { + hdr, _, _, err := ParsePacket(data) + if err != nil { + b.Fatal(err) + } + if _, err := hdr.ParseExtended(data); err != nil { + b.Fatal(err) + } + } +} diff --git a/third_party/quic-go/internal/wire/frame.go b/third_party/quic-go/internal/wire/frame.go new file mode 100644 index 0000000..7468b5a --- /dev/null +++ b/third_party/quic-go/internal/wire/frame.go @@ -0,0 +1,33 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" +) + +// A Frame in QUIC +type Frame interface { + Append(b []byte, version protocol.Version) ([]byte, error) + Length(version protocol.Version) protocol.ByteCount +} + +// IsProbingFrame returns true if the frame is a probing frame. +// See section 9.1 of RFC 9000. +func IsProbingFrame(f Frame) bool { + switch f.(type) { + case *PathChallengeFrame, *PathResponseFrame, *NewConnectionIDFrame: + return true + } + return false +} + +// IsProbingFrameType returns true if the FrameType is a probing frame. +// See section 9.1 of RFC 9000. +func IsProbingFrameType(f FrameType) bool { + //nolint:exhaustive // PATH_CHALLENGE, PATH_RESPONSE and NEW_CONNECTION_ID are the only probing frames + switch f { + case FrameTypePathChallenge, FrameTypePathResponse, FrameTypeNewConnectionID: + return true + default: + return false + } +} diff --git a/third_party/quic-go/internal/wire/frame_parser.go b/third_party/quic-go/internal/wire/frame_parser.go new file mode 100644 index 0000000..8b6ef0c --- /dev/null +++ b/third_party/quic-go/internal/wire/frame_parser.go @@ -0,0 +1,192 @@ +package wire + +import ( + "errors" + "fmt" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/quicvarint" +) + +var errUnknownFrameType = errors.New("unknown frame type") + +// The FrameParser parses QUIC frames, one by one. +type FrameParser struct { + ackDelayExponent uint8 + supportsDatagrams bool + supportsResetStreamAt bool + supportsAckFrequency bool + + // To avoid allocating when parsing, keep a single ACK frame struct. + // It is used over and over again. + ackFrame *AckFrame +} + +// NewFrameParser creates a new frame parser. +func NewFrameParser(supportsDatagrams, supportsResetStreamAt, supportsAckFrequency bool) *FrameParser { + return &FrameParser{ + supportsDatagrams: supportsDatagrams, + supportsResetStreamAt: supportsResetStreamAt, + supportsAckFrequency: supportsAckFrequency, + ackFrame: &AckFrame{}, + } +} + +// ParseType parses the frame type of the next frame. +// It skips over PADDING frames. +func (p *FrameParser) ParseType(b []byte, encLevel protocol.EncryptionLevel) (FrameType, int, error) { + var parsed int + for len(b) != 0 { + typ, l, err := quicvarint.Parse(b) + parsed += l + if err != nil { + return 0, parsed, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + ErrorMessage: err.Error(), + } + } + b = b[l:] + if typ == 0x0 { // skip PADDING frames + continue + } + ft := FrameType(typ) + valid := ft.isValidRFC9000() || + (p.supportsDatagrams && ft.IsDatagramFrameType()) || + (p.supportsResetStreamAt && ft == FrameTypeResetStreamAt) || + (p.supportsAckFrequency && (ft == FrameTypeAckFrequency || ft == FrameTypeImmediateAck)) + if !valid { + return 0, parsed, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + FrameType: typ, + ErrorMessage: errUnknownFrameType.Error(), + } + } + if !ft.isAllowedAtEncLevel(encLevel) { + return 0, parsed, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + FrameType: typ, + ErrorMessage: fmt.Sprintf("%d not allowed at encryption level %s", ft, encLevel), + } + } + return ft, parsed, nil + } + return 0, parsed, io.EOF +} + +func (p *FrameParser) ParseStreamFrame(frameType FrameType, data []byte, v protocol.Version) (*StreamFrame, int, error) { + frame, n, err := ParseStreamFrame(data, frameType, v) + if err != nil { + return nil, n, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + FrameType: uint64(frameType), + ErrorMessage: err.Error(), + } + } + return frame, n, nil +} + +func (p *FrameParser) ParseAckFrame(frameType FrameType, data []byte, encLevel protocol.EncryptionLevel, v protocol.Version) (*AckFrame, int, error) { + ackDelayExponent := p.ackDelayExponent + if encLevel != protocol.Encryption1RTT { + ackDelayExponent = protocol.DefaultAckDelayExponent + } + p.ackFrame.Reset() + l, err := parseAckFrame(p.ackFrame, data, frameType, ackDelayExponent, v) + if err != nil { + return nil, l, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + FrameType: uint64(frameType), + ErrorMessage: err.Error(), + } + } + + return p.ackFrame, l, nil +} + +func (p *FrameParser) ParseDatagramFrame(frameType FrameType, data []byte, v protocol.Version) (*DatagramFrame, int, error) { + f, l, err := parseDatagramFrame(data, frameType, v) + if err != nil { + return nil, 0, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + FrameType: uint64(frameType), + ErrorMessage: err.Error(), + } + } + return f, l, nil +} + +// ParseLessCommonFrame parses everything except STREAM, ACK or DATAGRAM. +// These cases should be handled separately for performance reasons. +func (p *FrameParser) ParseLessCommonFrame(frameType FrameType, data []byte, v protocol.Version) (Frame, int, error) { + var frame Frame + var l int + var err error + //nolint:exhaustive // Common frames should already be handled. + switch frameType { + case FrameTypePing: + frame = &PingFrame{} + case FrameTypeResetStream: + frame, l, err = parseResetStreamFrame(data, false, v) + case FrameTypeStopSending: + frame, l, err = parseStopSendingFrame(data, v) + case FrameTypeCrypto: + frame, l, err = parseCryptoFrame(data, v) + case FrameTypeNewToken: + frame, l, err = parseNewTokenFrame(data, v) + case FrameTypeMaxData: + frame, l, err = parseMaxDataFrame(data, v) + case FrameTypeMaxStreamData: + frame, l, err = parseMaxStreamDataFrame(data, v) + case FrameTypeBidiMaxStreams, FrameTypeUniMaxStreams: + frame, l, err = parseMaxStreamsFrame(data, frameType, v) + case FrameTypeDataBlocked: + frame, l, err = parseDataBlockedFrame(data, v) + case FrameTypeStreamDataBlocked: + frame, l, err = parseStreamDataBlockedFrame(data, v) + case FrameTypeBidiStreamBlocked, FrameTypeUniStreamBlocked: + frame, l, err = parseStreamsBlockedFrame(data, frameType, v) + case FrameTypeNewConnectionID: + frame, l, err = parseNewConnectionIDFrame(data, v) + case FrameTypeRetireConnectionID: + frame, l, err = parseRetireConnectionIDFrame(data, v) + case FrameTypePathChallenge: + frame, l, err = parsePathChallengeFrame(data, v) + case FrameTypePathResponse: + frame, l, err = parsePathResponseFrame(data, v) + case FrameTypeConnectionClose, FrameTypeApplicationClose: + frame, l, err = parseConnectionCloseFrame(data, frameType, v) + case FrameTypeHandshakeDone: + frame = &HandshakeDoneFrame{} + case FrameTypeResetStreamAt: + frame, l, err = parseResetStreamFrame(data, true, v) + case FrameTypeAckFrequency: + frame, l, err = parseAckFrequencyFrame(data, v) + case FrameTypeImmediateAck: + frame = &ImmediateAckFrame{} + default: + err = errUnknownFrameType + } + if err != nil { + return frame, l, &qerr.TransportError{ + ErrorCode: qerr.FrameEncodingError, + FrameType: uint64(frameType), + ErrorMessage: err.Error(), + } + } + return frame, l, err +} + +// SetAckDelayExponent sets the acknowledgment delay exponent (sent in the transport parameters). +// This value is used to scale the ACK Delay field in the ACK frame. +func (p *FrameParser) SetAckDelayExponent(exp uint8) { + p.ackDelayExponent = exp +} + +func replaceUnexpectedEOF(e error) error { + if e == io.ErrUnexpectedEOF { + return io.EOF + } + return e +} diff --git a/third_party/quic-go/internal/wire/frame_parser_test.go b/third_party/quic-go/internal/wire/frame_parser_test.go new file mode 100644 index 0000000..8b16254 --- /dev/null +++ b/third_party/quic-go/internal/wire/frame_parser_test.go @@ -0,0 +1,1083 @@ +package wire + +import ( + "bytes" + "crypto/rand" + "fmt" + "io" + "slices" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/quicvarint" + + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" + + "github.com/stretchr/testify/require" +) + +func TestFrameTypeParsingReturnsNilWhenNothingToRead(t *testing.T) { + parser := NewFrameParser(true, true, true) + frameType, l, err := parser.ParseType(nil, protocol.Encryption1RTT) + require.Equal(t, io.EOF, err) + require.Zero(t, frameType) + require.Zero(t, l) +} + +func TestParseLessCommonFrameReturnsEOFWhenNothingToRead(t *testing.T) { + parser := NewFrameParser(true, true, true) + l, f, err := parser.ParseLessCommonFrame(FrameTypeMaxStreamData, nil, protocol.Version1) + require.IsType(t, &qerr.TransportError{}, err) + require.Zero(t, l) + require.Zero(t, f) +} + +func TestFrameParsingSkipsPaddingFrames(t *testing.T) { + parser := NewFrameParser(true, true, true) + b := []byte{0, 0} // 2 PADDING frames + b, err := (&PingFrame{}).Append(b, protocol.Version1) + require.NoError(t, err) + + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, 3, l) + require.Equal(t, FrameTypePing, frameType) + + frame, l, err := parser.ParseLessCommonFrame(frameType, b[1:], protocol.Version1) + require.NoError(t, err) + require.Zero(t, l) + require.IsType(t, &PingFrame{}, frame) +} + +func TestFrameParsingHandlesPaddingAtEnd(t *testing.T) { + parser := NewFrameParser(true, true, true) + b := []byte{0, 0, 0} + + _, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.Equal(t, io.EOF, err) + require.Equal(t, 3, l) +} + +func TestFrameParsingParsesSingleFrame(t *testing.T) { + parser := NewFrameParser(true, true, true) + var b []byte + for range 10 { + var err error + b, err = (&PingFrame{}).Append(b, protocol.Version1) + require.NoError(t, err) + } + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, FrameTypePing, frameType) + require.Equal(t, 1, l) + + frame, l, err := parser.ParseLessCommonFrame(frameType, b, protocol.Version1) + require.NoError(t, err) + require.Zero(t, l) + require.IsType(t, &PingFrame{}, frame) +} + +func TestFrameParserACK(t *testing.T) { + parser := NewFrameParser(true, true, true) + f := &AckFrame{AckRanges: []AckRange{{Smallest: 1, Largest: 0x13}}} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, FrameTypeAck, frameType) + require.Equal(t, 1, l) + + frame, l, err := parser.ParseAckFrame(frameType, b[l:], protocol.Encryption1RTT, protocol.Version1) + require.NoError(t, err) + require.NotNil(t, frame) + require.Equal(t, protocol.PacketNumber(0x13), frame.LargestAcked()) + require.Equal(t, len(b)-1, l) +} + +func TestFrameParserAckDelay(t *testing.T) { + t.Run("1-RTT", func(t *testing.T) { + testFrameParserAckDelay(t, protocol.Encryption1RTT) + }) + t.Run("Handshake", func(t *testing.T) { + testFrameParserAckDelay(t, protocol.EncryptionHandshake) + }) +} + +func testFrameParserAckDelay(t *testing.T, encLevel protocol.EncryptionLevel) { + parser := NewFrameParser(true, true, true) + parser.SetAckDelayExponent(protocol.AckDelayExponent + 2) + f := &AckFrame{ + AckRanges: []AckRange{{Smallest: 1, Largest: 1}}, + DelayTime: time.Second, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + frameType, l, err := parser.ParseType(b, encLevel) + require.NoError(t, err) + require.Equal(t, FrameTypeAck, frameType) + require.Equal(t, 1, l) + + frame, l, err := parser.ParseAckFrame(frameType, b[l:], encLevel, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(b)-1, l) + if encLevel == protocol.Encryption1RTT { + require.Equal(t, 4*time.Second, frame.DelayTime) + } else { + require.Equal(t, time.Second, frame.DelayTime) + } +} + +func checkFrameUnsupported(t *testing.T, err error, expectedFrameType uint64) { + t.Helper() + require.ErrorContains(t, err, errUnknownFrameType.Error()) + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.FrameEncodingError, transportErr.ErrorCode) + require.Equal(t, expectedFrameType, transportErr.FrameType) + require.Equal(t, "unknown frame type", transportErr.ErrorMessage) +} + +func TestFrameParserStreamFrames(t *testing.T) { + parser := NewFrameParser(true, true, true) + f := &StreamFrame{ + StreamID: 0x42, + Offset: 0x1337, + Fin: true, + Data: []byte("foobar"), + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, FrameType(0xd), frameType) + require.True(t, frameType.IsStreamFrameType()) + require.Equal(t, 1, l) + + // ParseLessCommonFrame should not handle Stream Frames + frame, l, err := parser.ParseLessCommonFrame(frameType, b[l:], protocol.Version1) + checkFrameUnsupported(t, err, 0xd) + require.Nil(t, frame) + require.Zero(t, l) +} + +func TestParseStreamFrameWrapsError(t *testing.T) { + parser := NewFrameParser(true, true, true) + f := &StreamFrame{ + StreamID: 0x1234, + Offset: 0x1000, + Data: []byte("hello world"), + DataLenPresent: true, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + + // Corrupt the buffer to trigger a parse error + b = b[:len(b)-2] // Remove last 2 bytes to cause an EOF + + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + + frame, n, err := parser.ParseStreamFrame(frameType, b[l:], protocol.Version1) + require.Nil(t, frame) + require.Zero(t, n) + + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.FrameEncodingError, transportErr.ErrorCode) + require.Equal(t, uint64(frameType), transportErr.FrameType) + require.Contains(t, transportErr.Error(), "EOF") +} + +func TestParseStreamFrameSuccess(t *testing.T) { + parser := NewFrameParser(true, true, true) + original := &StreamFrame{ + StreamID: 0x1234, + Offset: 0x1000, + Fin: true, + Data: []byte("hello world"), + DataLenPresent: true, + } + b, err := original.Append(nil, protocol.Version1) + require.NoError(t, err) + + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.True(t, frameType.IsStreamFrameType()) + require.Equal(t, FrameType(0x0f), frameType) // STREAM | OFF | LEN | FIN + + parsed, n, err := parser.ParseStreamFrame(frameType, b[l:], protocol.Version1) + require.NoError(t, err) + require.NotNil(t, parsed) + require.Equal(t, len(b)-l, n) + + require.Equal(t, original.StreamID, parsed.StreamID) + require.Equal(t, original.Offset, parsed.Offset) + require.Equal(t, original.Fin, parsed.Fin) + require.Equal(t, original.DataLenPresent, parsed.DataLenPresent) + require.Equal(t, original.Data, parsed.Data) +} + +func TestFrameParserFrames(t *testing.T) { + tests := []struct { + name string + frameType FrameType + frame Frame + }{ + { + name: "MAX_DATA", + frameType: FrameTypeMaxData, + frame: &MaxDataFrame{MaximumData: 0xcafe}, + }, + { + name: "MAX_STREAM_DATA", + frameType: FrameTypeMaxStreamData, + frame: &MaxStreamDataFrame{StreamID: 0xdeadbeef, MaximumStreamData: 0xdecafbad}, + }, + { + name: "RESET_STREAM", + frameType: FrameTypeResetStream, + frame: &ResetStreamFrame{ + StreamID: 0xdeadbeef, + FinalSize: 0xdecafbad1234, + ErrorCode: 0x1337, + }, + }, + { + name: "STOP_SENDING", + frameType: FrameTypeStopSending, + frame: &StopSendingFrame{StreamID: 0x42}, + }, + { + name: "CRYPTO", + frameType: FrameTypeCrypto, + frame: &CryptoFrame{Offset: 0x1337, Data: []byte("lorem ipsum")}, + }, + { + name: "NEW_TOKEN", + frameType: FrameTypeNewToken, + frame: &NewTokenFrame{Token: []byte("foobar")}, + }, + { + name: "MAX_STREAMS", + frameType: FrameTypeBidiMaxStreams, + frame: &MaxStreamsFrame{Type: protocol.StreamTypeBidi, MaxStreamNum: 0x1337}, + }, + { + name: "DATA_BLOCKED", + frameType: FrameTypeDataBlocked, + frame: &DataBlockedFrame{MaximumData: 0x1234}, + }, + { + name: "STREAM_DATA_BLOCKED", + frameType: FrameTypeStreamDataBlocked, + frame: &StreamDataBlockedFrame{StreamID: 0xdeadbeef, MaximumStreamData: 0xdead}, + }, + { + name: "STREAMS_BLOCKED", + frameType: FrameTypeBidiStreamBlocked, + frame: &StreamsBlockedFrame{Type: protocol.StreamTypeBidi, StreamLimit: 0x1234567}, + }, + { + name: "NEW_CONNECTION_ID", + frameType: FrameTypeNewConnectionID, + frame: &NewConnectionIDFrame{ + SequenceNumber: 0x1337, + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15}, + }, + }, + { + name: "RETIRE_CONNECTION_ID", + frameType: FrameTypeRetireConnectionID, + frame: &RetireConnectionIDFrame{SequenceNumber: 0x1337}, + }, + { + name: "PATH_CHALLENGE", + frameType: FrameTypePathChallenge, + frame: &PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + }, + { + name: "PATH_RESPONSE", + frameType: FrameTypePathResponse, + frame: &PathResponseFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + }, + { + name: "CONNECTION_CLOSE", + frameType: FrameTypeConnectionClose, + frame: &ConnectionCloseFrame{IsApplicationError: false, ReasonPhrase: "foobar"}, + }, + { + name: "APPLICATION_CLOSE", + frameType: FrameTypeApplicationClose, + frame: &ConnectionCloseFrame{IsApplicationError: true, ReasonPhrase: "foobar"}, + }, + { + name: "HANDSHAKE_DONE", + frameType: FrameTypeHandshakeDone, + frame: &HandshakeDoneFrame{}, + }, + { + name: "RESET_STREAM_AT", + frameType: FrameTypeResetStreamAt, + frame: &ResetStreamFrame{StreamID: 0x1337, ReliableSize: 0x42, FinalSize: 0xdeadbeef}, + }, + { + name: "ACK_FREQUENCY", + frameType: FrameTypeAckFrequency, + frame: &AckFrequencyFrame{ + SequenceNumber: 0x1337, + AckElicitingThreshold: 0x42, + RequestMaxAckDelay: 123 * time.Second, + ReorderingThreshold: 0xcafe, + }, + }, + { + name: "IMMEDIATE_ACK", + frameType: FrameTypeImmediateAck, + frame: &ImmediateAckFrame{}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + parser := NewFrameParser(true, true, true) + b, err := test.frame.Append(nil, protocol.Version1) + require.NoError(t, err) + + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, test.frameType, frameType) + require.Equal(t, quicvarint.Len(uint64(test.frameType)), l) + + frame, l, err := parser.ParseLessCommonFrame(frameType, b[l:], protocol.Version1) + require.NoError(t, err) + require.Equal(t, test.frame, frame) + require.Equal(t, len(b)-quicvarint.Len(uint64(test.frameType)), l) + }) + } +} + +func TestFrameAllowedAtEncLevel(t *testing.T) { + type testCase struct { + name string + frameType FrameType + frame Frame + allowedInitial bool + allowedHandshake bool + allowedZeroRTT bool + allowedOneRTT bool + } + + for _, tc := range []testCase{ + { + name: "CRYPTO_FRAME", + frameType: FrameTypeCrypto, + frame: &CryptoFrame{Offset: 0, Data: []byte("foo")}, + allowedInitial: true, + allowedHandshake: true, + allowedZeroRTT: false, + allowedOneRTT: true, + }, + { + name: "ACK_FRAME", + frameType: FrameTypeAck, + frame: &AckFrame{AckRanges: []AckRange{{Smallest: 1, Largest: 1}}}, + allowedInitial: true, + allowedHandshake: true, + allowedZeroRTT: false, + allowedOneRTT: true, + }, + { + name: "CONNECTION_CLOSE_FRAME", + frameType: FrameTypeConnectionClose, + frame: &ConnectionCloseFrame{IsApplicationError: false, ReasonPhrase: "err"}, + allowedInitial: true, + allowedHandshake: true, + allowedZeroRTT: false, + allowedOneRTT: true, + }, + { + name: "PING_FRAME", + frameType: FrameTypePing, + frame: &PingFrame{}, + allowedInitial: true, + allowedHandshake: true, + allowedZeroRTT: true, + allowedOneRTT: true, + }, + { + name: "NEW_TOKEN_FRAME", + frameType: FrameTypeNewToken, + frame: &NewTokenFrame{Token: []byte("tok")}, + allowedInitial: false, + allowedHandshake: false, + allowedZeroRTT: false, + allowedOneRTT: true, + }, + { + name: "PATH_RESPONSE_FRAME", + frameType: FrameTypePathResponse, + frame: &PathResponseFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + allowedInitial: false, + allowedHandshake: false, + allowedZeroRTT: false, + allowedOneRTT: true, + }, + { + name: "RETIRE_CONNECTION_ID_FRAME", + frameType: FrameTypeRetireConnectionID, + frame: &RetireConnectionIDFrame{SequenceNumber: 1}, + allowedInitial: false, + allowedHandshake: false, + allowedZeroRTT: false, + allowedOneRTT: true, + }, + { + name: "MAX_DATA_FRAME", + frameType: FrameTypeMaxData, + frame: &MaxDataFrame{MaximumData: 1}, + allowedInitial: false, + allowedHandshake: false, + allowedZeroRTT: true, + allowedOneRTT: true, + }, + { + name: "STREAM_FRAME", + frameType: FrameType(0x8), + frame: &StreamFrame{StreamID: 1, Data: []byte("foobar")}, + allowedInitial: false, + allowedHandshake: false, + allowedZeroRTT: true, + allowedOneRTT: true, + }, + { + name: "RESET_STREAM_AT", + frameType: FrameTypeResetStreamAt, + frame: &ResetStreamFrame{StreamID: 1, FinalSize: 1, ReliableSize: 1}, + allowedInitial: false, + allowedHandshake: false, + allowedZeroRTT: true, + allowedOneRTT: true, + }, + } { + for _, encLevel := range []protocol.EncryptionLevel{ + protocol.EncryptionInitial, + protocol.EncryptionHandshake, + protocol.Encryption0RTT, + protocol.Encryption1RTT, + } { + t.Run(fmt.Sprintf("%s/%v", tc.name, encLevel), func(t *testing.T) { + var allowed bool + switch encLevel { + case protocol.EncryptionInitial: + allowed = tc.allowedInitial + case protocol.EncryptionHandshake: + allowed = tc.allowedHandshake + case protocol.Encryption0RTT: + allowed = tc.allowedZeroRTT + case protocol.Encryption1RTT: + allowed = tc.allowedOneRTT + } + + parser := NewFrameParser(true, true, true) + b, err := tc.frame.Append(nil, protocol.Version1) + require.NoError(t, err) + frameType, _, err := parser.ParseType(b, encLevel) + if allowed { + require.NoError(t, err) + require.Equal(t, tc.frameType, frameType) + } else { + require.Error(t, err) + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.FrameEncodingError, transportErr.ErrorCode) + } + }) + } + } +} + +func TestFrameParserDatagramFrame(t *testing.T) { + parser := NewFrameParser(true, true, true) + f := &DatagramFrame{ + Data: []byte("foobar"), + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + frameType, l, err := parser.ParseType(b, protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, FrameTypeDatagramNoLength, frameType) + require.Equal(t, 1, l) + + // ParseLessCommonFrame should not be used to handle DATAGRAM frames + _, _, err = parser.ParseLessCommonFrame(frameType, b[l:], protocol.Version1) + require.Error(t, err) + + // parseDatagramFrame should be used for this type + datagramFrame, l, err := parser.ParseDatagramFrame(frameType, b[l:], protocol.Version1) + require.NoError(t, err) + require.IsType(t, &DatagramFrame{}, datagramFrame) + require.Equal(t, 6, l) + require.Equal(t, f.Data, datagramFrame.Data) +} + +func TestFrameParserDatagramUnsupported(t *testing.T) { + parser := NewFrameParser(false, true, true) + f := &DatagramFrame{Data: []byte("foobar")} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + + _, _, err = parser.ParseType(b, protocol.Encryption1RTT) + checkFrameUnsupported(t, err, 0x30) +} + +func TestFrameParserResetStreamAtUnsupported(t *testing.T) { + parser := NewFrameParser(true, false, true) + f := &ResetStreamFrame{StreamID: 0x1337, ReliableSize: 0x42, FinalSize: 0xdeadbeef} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + + _, _, err = parser.ParseType(b, protocol.Encryption1RTT) + checkFrameUnsupported(t, err, uint64(FrameTypeResetStreamAt)) +} + +func TestFrameParserAckFrequencyUnsupported(t *testing.T) { + parser := NewFrameParser(true, true, false) + + t.Run("ACK_FREQUENCY", func(t *testing.T) { + f := &AckFrequencyFrame{ + SequenceNumber: 1337, + AckElicitingThreshold: 42, + RequestMaxAckDelay: 42 * time.Millisecond, + ReorderingThreshold: 1234, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + _, _, err = parser.ParseType(b, protocol.Encryption1RTT) + checkFrameUnsupported(t, err, uint64(FrameTypeAckFrequency)) + }) + + t.Run("IMMEDIATE_ACK", func(t *testing.T) { + f := &ImmediateAckFrame{} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + _, _, err = parser.ParseType(b, protocol.Encryption1RTT) + checkFrameUnsupported(t, err, uint64(FrameTypeImmediateAck)) + }) +} + +func TestFrameParserInvalidFrameType(t *testing.T) { + parser := NewFrameParser(true, true, true) + + _, l, err := parser.ParseType(encodeVarInt(0x42), protocol.Encryption1RTT) + + require.Equal(t, 2, l) + + require.Error(t, err) + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.FrameEncodingError, transportErr.ErrorCode) +} + +func TestFrameParsingErrorsOnInvalidFrames(t *testing.T) { + parser := NewFrameParser(true, true, true) + f := &MaxStreamDataFrame{ + StreamID: 0x1337, + MaximumStreamData: 0xdeadbeef, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + + frameType, l, err := parser.ParseType(b[:len(b)-2], protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, FrameTypeMaxStreamData, frameType) + require.Equal(t, 1, l) + + _, _, err = parser.ParseLessCommonFrame(frameType, b[1:len(b)-2], protocol.Version1) + require.Error(t, err) + var transportErr *qerr.TransportError + require.ErrorAs(t, err, &transportErr) + require.Equal(t, qerr.FrameEncodingError, transportErr.ErrorCode) +} + +func writeFrames(tb testing.TB, frames ...Frame) []byte { + var b []byte + for _, f := range frames { + var err error + b, err = f.Append(b, protocol.Version1) + require.NoError(tb, err) + } + return b +} + +// This function is used in benchmarks, and also to ensure zero allocation for STREAM frame parsing. +// We can therefore not use the require framework, as it allocates. +func parseFrames(tb testing.TB, parser *FrameParser, data []byte, frames ...Frame) { + for _, expectedFrame := range frames { + frameType, l, err := parser.ParseType(data, protocol.Encryption1RTT) + if err != nil { + tb.Fatal(err) + } + data = data[l:] + + if frameType.IsStreamFrameType() { + sf := expectedFrame.(*StreamFrame) + frame, l, err := ParseStreamFrame(data, frameType, protocol.Version1) + if err != nil { + tb.Fatal(err) + } + if sf.StreamID != frame.StreamID || sf.Offset != frame.Offset { + tb.Fatalf("STREAM frame does not match: %v vs %v", sf, frame) + } + frame.PutBack() + data = data[l:] + continue + } + + if frameType.IsAckFrameType() { + af, ok := expectedFrame.(*AckFrame) + if !ok { + tb.Fatalf("expected ACK, but got %v", expectedFrame) + } + + f, l, err := parser.ParseAckFrame(frameType, data, protocol.Encryption1RTT, protocol.Version1) + if f.DelayTime != af.DelayTime || f.ECNCE != af.ECNCE || f.ECT0 != af.ECT0 || f.ECT1 != af.ECT1 { + tb.Fatal(err) + } + if f.DelayTime != af.DelayTime { + tb.Fatalf("ACK frame does not match: %v vs %v", af, f) + } + if !slices.Equal(f.AckRanges, af.AckRanges) { + tb.Fatalf("ACK frame ACK ranges don't match: %v vs %v", af, f) + } + data = data[l:] + continue + } + + if frameType.IsDatagramFrameType() { + df, ok := expectedFrame.(*DatagramFrame) + if !ok { + tb.Fatalf("expected DATAGRAM, but got %v", expectedFrame) + } + + f, l, err := parser.ParseDatagramFrame(frameType, data, protocol.Version1) + if err != nil { + tb.Fatal(err) + } + if df.DataLenPresent != f.DataLenPresent || !bytes.Equal(df.Data, f.Data) { + tb.Fatalf("DATAGRAM frame does not match: %v vs %v", df, f) + } + data = data[l:] + continue + } + + f, l, err := parser.ParseLessCommonFrame(frameType, data, protocol.Version1) + if err != nil { + tb.Fatal(err) + } + data = data[l:] + + switch frameType { + case FrameTypeMaxData: + mdf, ok := expectedFrame.(*MaxDataFrame) + if !ok { + tb.Fatalf("expected MAX_DATA, but got %v", expectedFrame) + } + if *f.(*MaxDataFrame) != *mdf { + tb.Fatalf("MAX_DATA frame does not match: %v vs %v", f, mdf) + } + case FrameTypeUniMaxStreams: + msf, ok := expectedFrame.(*MaxStreamsFrame) + if !ok { + tb.Fatalf("expected MAX_STREAMS, but got %v", expectedFrame) + } + if *f.(*MaxStreamsFrame) != *msf { + tb.Fatalf("MAX_STREAMS frame does not match: %v vs %v", f, msf) + } + case FrameTypeMaxStreamData: + mdf, ok := expectedFrame.(*MaxStreamDataFrame) + if !ok { + tb.Fatalf("expected MAX_STREAM_DATA, but got %v", expectedFrame) + } + if *f.(*MaxStreamDataFrame) != *mdf { + tb.Fatalf("MAX_STREAM_DATA frame does not match: %v vs %v", f, mdf) + } + case FrameTypeCrypto: + cf, ok := expectedFrame.(*CryptoFrame) + if !ok { + tb.Fatalf("expected CRYPTO, but got %v", expectedFrame) + } + frame := f.(*CryptoFrame) + if frame.Offset != cf.Offset || !bytes.Equal(frame.Data, cf.Data) { + tb.Fatalf("CRYPTO frame does not match: %v vs %v", f, cf) + } + case FrameTypePing: + _ = f.(*PingFrame) + case FrameTypeResetStream: + rsf, ok := expectedFrame.(*ResetStreamFrame) + if !ok { + tb.Fatalf("expected RESET_STREAM, but got %v", expectedFrame) + } + if *f.(*ResetStreamFrame) != *rsf { + tb.Fatalf("RESET_STREAM frame does not match: %v vs %v", f, rsf) + } + continue + default: + tb.Fatalf("Frame type not supported in benchmark or should not occur: %v", frameType) + } + } +} + +func TestFrameParserAllocs(t *testing.T) { + t.Run("STREAM", func(t *testing.T) { + var frames []Frame + for i := range 10 { + frames = append(frames, &StreamFrame{ + StreamID: protocol.StreamID(1337 + i), + Offset: protocol.ByteCount(1e7 + i), + Data: make([]byte, 200+i), + DataLenPresent: true, + }) + } + require.Zero(t, testFrameParserAllocs(t, frames)) + }) + + t.Run("ACK", func(t *testing.T) { + var frames []Frame + for i := range 10 { + frames = append(frames, &AckFrame{ + AckRanges: []AckRange{ + {Smallest: protocol.PacketNumber(5000 + i), Largest: protocol.PacketNumber(5200 + i)}, + {Smallest: protocol.PacketNumber(1 + i), Largest: protocol.PacketNumber(4200 + i)}, + }, + DelayTime: time.Duration(int64(time.Millisecond) * int64(i)), + ECT0: uint64(5000 + i), + ECT1: uint64(i), + ECNCE: uint64(10 + i), + }) + } + require.Zero(t, testFrameParserAllocs(t, frames)) + }) +} + +func testFrameParserAllocs(t *testing.T, frames []Frame) float64 { + buf := writeFrames(t, frames...) + parser := NewFrameParser(true, true, true) + parser.SetAckDelayExponent(3) + + return testing.AllocsPerRun(100, func() { + parseFrames(t, parser, buf, frames...) + }) +} + +func BenchmarkParseOtherFrames(b *testing.B) { + frames := []Frame{ + &MaxDataFrame{MaximumData: 123456}, + &MaxStreamsFrame{MaxStreamNum: 10}, + &MaxStreamDataFrame{StreamID: 1337, MaximumStreamData: 1e6}, + &CryptoFrame{Offset: 1000, Data: make([]byte, 128)}, + &PingFrame{}, + &ResetStreamFrame{StreamID: 87654, ErrorCode: 1234, FinalSize: 1e8}, + } + benchmarkFrames(b, frames...) +} + +func BenchmarkParseAckFrame(b *testing.B) { + var frames []Frame + for i := range 10 { + frames = append(frames, &AckFrame{ + AckRanges: []AckRange{ + {Smallest: protocol.PacketNumber(5000 + i), Largest: protocol.PacketNumber(5200 + i)}, + {Smallest: protocol.PacketNumber(1 + i), Largest: protocol.PacketNumber(4200 + i)}, + }, + DelayTime: time.Duration(int64(time.Millisecond) * int64(i)), + ECT0: uint64(5000 + i), + ECT1: uint64(i), + ECNCE: uint64(10 + i), + }) + } + benchmarkFrames(b, frames...) +} + +func BenchmarkParseStreamFrame(b *testing.B) { + var frames []Frame + for i := range 10 { + data := make([]byte, 200+i) + rand.Read(data) + frames = append(frames, &StreamFrame{ + StreamID: protocol.StreamID(1337 + i), + Offset: protocol.ByteCount(1e7 + i), + Data: data, + DataLenPresent: true, + }) + } + benchmarkFrames(b, frames...) +} + +func BenchmarkParseDatagramFrame(b *testing.B) { + var frames []Frame + for i := range 10 { + data := make([]byte, 200+i) + rand.Read(data) + frames = append(frames, &DatagramFrame{ + Data: data, + DataLenPresent: true, + }) + } + benchmarkFrames(b, frames...) +} + +func benchmarkFrames(b *testing.B, frames ...Frame) { + b.ReportAllocs() + + buf := writeFrames(b, frames...) + parser := NewFrameParser(true, true, true) + parser.SetAckDelayExponent(3) + + for b.Loop() { + parseFrames(b, parser, buf, frames...) + } +} + +func FuzzFrames(f *testing.F) { + corpus := ossfuzzseeds.New(f) + + const version = protocol.Version1 + + for _, s := range []struct { + encLevel protocol.EncryptionLevel + frame Frame + }{ + {encLevel: protocol.EncryptionInitial, frame: &PingFrame{}}, + {encLevel: protocol.EncryptionHandshake, frame: &PingFrame{}}, + {encLevel: protocol.Encryption0RTT, frame: &PingFrame{}}, + {encLevel: protocol.EncryptionInitial, frame: &CryptoFrame{Offset: 42, Data: []byte("initial crypto")}}, + {encLevel: protocol.EncryptionHandshake, frame: &CryptoFrame{Offset: 123, Data: []byte("handshake crypto")}}, + {encLevel: protocol.EncryptionInitial, frame: &AckFrame{AckRanges: []AckRange{{Smallest: 1, Largest: 10}}}}, + {encLevel: protocol.EncryptionHandshake, frame: &AckFrame{AckRanges: []AckRange{{Smallest: 1, Largest: 10}}}}, + } { + b, err := s.frame.Append(nil, version) + require.NoError(f, err) + corpus.Add(uint8(s.encLevel), uint16(protocol.MaxPacketBufferSize), b) + } + + for _, fr := range []Frame{ + &PingFrame{}, + &StreamFrame{StreamID: 0x42, Fin: true}, + &StreamFrame{StreamID: 0x42, Data: []byte("foobar"), Fin: true}, + &StreamFrame{StreamID: 0x1337, Offset: 0xcafe, Data: []byte("foobar")}, + &StreamFrame{Offset: quicvarint.Max, Data: []byte("foo")}, // exceeds maximum offset + &AckFrame{AckRanges: []AckRange{{Smallest: 1, Largest: 0x13}}}, + &AckFrame{ + AckRanges: []AckRange{{Smallest: 80, Largest: 100}, {Smallest: 1, Largest: 50}}, + DelayTime: time.Millisecond, + ECT0: 42, + ECT1: 13, + ECNCE: 7, + }, + &ResetStreamFrame{StreamID: 0x1337, ErrorCode: 0x42, FinalSize: 0xdead}, + &StopSendingFrame{StreamID: 0x42, ErrorCode: 0x1337}, + &CryptoFrame{Offset: 0x1337, Data: []byte("crypto data")}, + &NewTokenFrame{Token: []byte("token")}, + &MaxDataFrame{MaximumData: 0xcafe}, + &MaxStreamDataFrame{StreamID: 0xdead, MaximumStreamData: 0xbeef}, + &MaxStreamsFrame{Type: protocol.StreamTypeBidi, MaxStreamNum: 0x42}, + &MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: 0x42}, + &DataBlockedFrame{MaximumData: 0x1234}, + &StreamDataBlockedFrame{StreamID: 0xdead, MaximumStreamData: 0xbeef}, + &StreamsBlockedFrame{Type: protocol.StreamTypeBidi, StreamLimit: 0x42}, + &StreamsBlockedFrame{Type: protocol.StreamTypeUni, StreamLimit: 0x42}, + &NewConnectionIDFrame{ + SequenceNumber: 0x42, + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15}, + }, + &RetireConnectionIDFrame{SequenceNumber: 0x42}, + &PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + &PathResponseFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + &ConnectionCloseFrame{ErrorCode: 0x42, ReasonPhrase: "foobar"}, + &ConnectionCloseFrame{IsApplicationError: true, ErrorCode: 0x42, ReasonPhrase: "foobar"}, + &HandshakeDoneFrame{}, + &DatagramFrame{Data: []byte("datagram")}, + &ResetStreamFrame{StreamID: 0x1337, ReliableSize: 0x42, FinalSize: 0xdead}, + &AckFrequencyFrame{ + SequenceNumber: 0x42, + AckElicitingThreshold: 10, + RequestMaxAckDelay: 25 * time.Millisecond, + ReorderingThreshold: 5, + }, + &ImmediateAckFrame{}, + } { + b, err := fr.Append(nil, version) + require.NoError(f, err) + maxSize := uint16(protocol.MaxPacketBufferSize) + switch fr.(type) { + case *StreamFrame, *DatagramFrame: + maxSize = 256 + case *AckFrame: + maxSize = 128 + } + corpus.Add(uint8(protocol.Encryption1RTT), maxSize, b) + } + + f.Fuzz(func(t *testing.T, encLevelRaw uint8, maxSize uint16, data []byte) { + encLevel := protocol.EncryptionLevel(encLevelRaw) + if encLevel != protocol.EncryptionInitial && encLevel != protocol.EncryptionHandshake && encLevel != protocol.Encryption1RTT && encLevel != protocol.Encryption0RTT { + return + } + // maxSize is used to split off frames from the original frame (in the case of CRYPTO and STREAM frames), + // and to truncate ACK frames. + // This happens at the packet boundary, so values larger than the packet size are not interesting. + if maxSize > 10_000 { + return + } + + parser := NewFrameParser(true, true, true) + parser.SetAckDelayExponent(protocol.DefaultAckDelayExponent) + + var b []byte + for len(data) > 0 { + initialLen := len(data) + frameType, l, err := parser.ParseType(data, encLevel) + if err != nil { + return + } + data = data[l:] + + var frame Frame + switch { + case frameType.IsStreamFrameType(): + frame, l, err = parser.ParseStreamFrame(frameType, data, version) + case frameType.IsAckFrameType(): + frame, l, err = parser.ParseAckFrame(frameType, data, encLevel, version) + case frameType == FrameTypeDatagramNoLength || frameType == FrameTypeDatagramWithLength: + frame, l, err = parser.ParseDatagramFrame(frameType, data, version) + default: + frame, l, err = parser.ParseLessCommonFrame(frameType, data, version) + } + if err != nil { + return + } + data = data[l:] + IsProbingFrame(frame) + + if sf, ok := frame.(*StreamFrame); ok { + if sf.DataLen() == 0 { + sf.PutBack() + continue + } + } + checkFrameInvariants(t, frame) + + startLen := len(b) + parsedLen := initialLen - len(data) + b, err = frame.Append(b, version) + require.NoError(t, err) + frameLen := protocol.ByteCount(len(b) - startLen) + require.Equal(t, frameLen, frame.Length(version), "re-serialized frame") + size := protocol.ByteCount(maxSize) + switch f := frame.(type) { + case *StreamFrame: + orig := slices.Clone(f.Data) + split, needsSplit := f.MaybeSplitOffFrame(size, version) + if split != nil { + require.LessOrEqual(t, split.Length(version), size, "split STREAM frame") + require.Equal(t, orig, append(slices.Clone(split.Data), f.Data...), "split STREAM frame data") + split.PutBack() + } else { + require.True(t, needsSplit || f.Length(version) <= size, "STREAM longer than maxSize but not split: len=%d maxSize=%d", f.Length(version), size) + } + case *CryptoFrame: + orig := slices.Clone(f.Data) + split, needsSplit := f.MaybeSplitOffFrame(size, version) + if split != nil { + require.LessOrEqual(t, split.Length(version), size, "split CRYPTO frame") + require.Equal(t, orig, append(slices.Clone(split.Data), f.Data...), "split CRYPTO frame data") + } else { + require.True(t, needsSplit || f.Length(version) <= size, "CRYPTO longer than maxSize but not split: len=%d maxSize=%d", f.Length(version), size) + } + case *AckFrame: + f.HasMissingRanges() + // Truncate requires maxSize to fit at least one ACK range; + // 64 bytes covers the worst case (8-byte varints, with ECN). + if size >= 64 { + f.Truncate(size, version) + require.LessOrEqual(t, f.Length(version), size, "truncated ACK") + checkFrameInvariants(t, f) + } + case *DatagramFrame: + if n := f.MaxDataLen(size, version); n > 0 && protocol.ByteCount(len(f.Data)) > n { + orig := f.Data + f.Data = f.Data[:n] + require.LessOrEqual(t, f.Length(version), size, "DATAGRAM with MaxDataLen") + f.Data = orig + } + } + if sf, ok := frame.(*StreamFrame); ok { + sf.PutBack() + } + require.LessOrEqual(t, frameLen, protocol.ByteCount(parsedLen), "serialized length vs parsed length") + } + }) +} + +func checkFrameInvariants(t *testing.T, frame Frame) { + t.Helper() + + switch f := frame.(type) { + case *StreamFrame: + if protocol.ByteCount(len(f.Data)) != f.DataLen() { + t.Fatal("STREAM frame: inconsistent data length") + } + case *AckFrame: + if f.DelayTime < 0 { + t.Fatalf("invalid ACK delay_time: %s", f.DelayTime) + } + if f.LargestAcked() < f.LowestAcked() { + t.Fatal("ACK: largest acknowledged is smaller than lowest acknowledged") + } + for _, r := range f.AckRanges { + if r.Largest < 0 || r.Smallest < 0 { + t.Fatal("ACK range contains a negative packet number") + } + } + if !f.AcksPacket(f.LargestAcked()) { + t.Fatal("ACK frame claims that largest acknowledged is not acknowledged") + } + if !f.AcksPacket(f.LowestAcked()) { + t.Fatal("ACK frame claims that lowest acknowledged is not acknowledged") + } + _ = f.AcksPacket(100) + _ = f.AcksPacket((f.LargestAcked() + f.LowestAcked()) / 2) + case *NewConnectionIDFrame: + if f.ConnectionID.Len() < 1 || f.ConnectionID.Len() > 20 { + t.Fatalf("invalid NEW_CONNECTION_ID frame length: %s", f.ConnectionID) + } + case *NewTokenFrame: + if len(f.Token) == 0 { + t.Fatal("NEW_TOKEN frame with an empty token") + } + case *MaxStreamsFrame: + if f.MaxStreamNum > protocol.MaxStreamCount { + t.Fatal("MAX_STREAMS frame with an invalid Maximum Streams value") + } + case *StreamsBlockedFrame: + if f.StreamLimit > protocol.MaxStreamCount { + t.Fatal("STREAMS_BLOCKED frame with an invalid Maximum Streams value") + } + case *ConnectionCloseFrame: + if f.IsApplicationError && f.FrameType != 0 { + t.Fatal("CONNECTION_CLOSE for an application error containing a frame type") + } + case *ResetStreamFrame: + if f.FinalSize < f.ReliableSize { + t.Fatal("RESET_STREAM frame with a FinalSize smaller than the ReliableSize") + } + case *AckFrequencyFrame: + if f.RequestMaxAckDelay < 0 { + t.Fatal("ACK_FREQUENCY frame with a negative RequestMaxAckDelay") + } + } +} diff --git a/third_party/quic-go/internal/wire/frame_test.go b/third_party/quic-go/internal/wire/frame_test.go new file mode 100644 index 0000000..5dbfba8 --- /dev/null +++ b/third_party/quic-go/internal/wire/frame_test.go @@ -0,0 +1,42 @@ +package wire + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestProbingFrames(t *testing.T) { + testCases := map[Frame]bool{ + &AckFrame{}: false, + &ConnectionCloseFrame{}: false, + &DataBlockedFrame{}: false, + &PingFrame{}: false, + &ResetStreamFrame{}: false, + &StreamFrame{}: false, + &DatagramFrame{}: false, + &MaxDataFrame{}: false, + &MaxStreamDataFrame{}: false, + &StopSendingFrame{}: false, + &PathChallengeFrame{}: true, + &PathResponseFrame{}: true, + &NewConnectionIDFrame{}: true, + } + + for f, expected := range testCases { + require.Equal(t, expected, IsProbingFrame(f)) + } +} + +func TestIsProbingFrameType(t *testing.T) { + tests := map[FrameType]bool{ + FrameTypePathChallenge: true, + FrameTypePathResponse: true, + FrameTypeNewConnectionID: true, + FrameType(0x01): false, + FrameType(0xFF): false, + } + for ft, expected := range tests { + require.Equal(t, expected, IsProbingFrameType(ft)) + } +} diff --git a/third_party/quic-go/internal/wire/frame_type.go b/third_party/quic-go/internal/wire/frame_type.go new file mode 100644 index 0000000..84b01fc --- /dev/null +++ b/third_party/quic-go/internal/wire/frame_type.go @@ -0,0 +1,81 @@ +package wire + +import "github.com/apernet/quic-go/internal/protocol" + +type FrameType uint64 + +// These constants correspond to those defined in RFC 9000. +// Stream frame types are not listed explicitly here; use FrameType.IsStreamFrameType() to identify them. +const ( + FrameTypePing FrameType = 0x1 + FrameTypeAck FrameType = 0x2 + FrameTypeAckECN FrameType = 0x3 + FrameTypeResetStream FrameType = 0x4 + FrameTypeStopSending FrameType = 0x5 + FrameTypeCrypto FrameType = 0x6 + FrameTypeNewToken FrameType = 0x7 + + FrameTypeMaxData FrameType = 0x10 + FrameTypeMaxStreamData FrameType = 0x11 + FrameTypeBidiMaxStreams FrameType = 0x12 + FrameTypeUniMaxStreams FrameType = 0x13 + FrameTypeDataBlocked FrameType = 0x14 + FrameTypeStreamDataBlocked FrameType = 0x15 + FrameTypeBidiStreamBlocked FrameType = 0x16 + FrameTypeUniStreamBlocked FrameType = 0x17 + FrameTypeNewConnectionID FrameType = 0x18 + FrameTypeRetireConnectionID FrameType = 0x19 + FrameTypePathChallenge FrameType = 0x1a + FrameTypePathResponse FrameType = 0x1b + FrameTypeConnectionClose FrameType = 0x1c + FrameTypeApplicationClose FrameType = 0x1d + FrameTypeHandshakeDone FrameType = 0x1e + // https://datatracker.ietf.org/doc/draft-ietf-quic-reliable-stream-reset/09/ + FrameTypeResetStreamAt FrameType = 0x24 + // https://datatracker.ietf.org/doc/draft-ietf-quic-ack-frequency/11/ + FrameTypeAckFrequency FrameType = 0xaf + FrameTypeImmediateAck FrameType = 0x1f + + FrameTypeDatagramNoLength FrameType = 0x30 + FrameTypeDatagramWithLength FrameType = 0x31 +) + +func (t FrameType) IsStreamFrameType() bool { + return t >= 0x8 && t <= 0xf +} + +func (t FrameType) isValidRFC9000() bool { + return t <= 0x1e +} + +func (t FrameType) IsAckFrameType() bool { + return t == FrameTypeAck || t == FrameTypeAckECN +} + +func (t FrameType) IsDatagramFrameType() bool { + return t == FrameTypeDatagramNoLength || t == FrameTypeDatagramWithLength +} + +func (t FrameType) isAllowedAtEncLevel(encLevel protocol.EncryptionLevel) bool { + //nolint:exhaustive + switch encLevel { + case protocol.EncryptionInitial, protocol.EncryptionHandshake: + switch t { + case FrameTypeCrypto, FrameTypeAck, FrameTypeAckECN, FrameTypeConnectionClose, FrameTypePing: + return true + default: + return false + } + case protocol.Encryption0RTT: + switch t { + case FrameTypeCrypto, FrameTypeAck, FrameTypeAckECN, FrameTypeConnectionClose, FrameTypeNewToken, FrameTypePathResponse, FrameTypeRetireConnectionID: + return false + default: + return true + } + case protocol.Encryption1RTT: + return true + default: + panic("unknown encryption level") + } +} diff --git a/third_party/quic-go/internal/wire/frame_type_test.go b/third_party/quic-go/internal/wire/frame_type_test.go new file mode 100644 index 0000000..2702af3 --- /dev/null +++ b/third_party/quic-go/internal/wire/frame_type_test.go @@ -0,0 +1,29 @@ +package wire + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestIsStreamFrameType(t *testing.T) { + for i := 0x08; i <= 0x0f; i++ { + require.Truef(t, FrameType(i).IsStreamFrameType(), "FrameType(0x%x).IsStreamFrameType() = false, want true", i) + } + + require.False(t, FrameType(0x1).IsStreamFrameType()) +} + +func TestIsAckFrameType(t *testing.T) { + require.True(t, FrameTypeAck.IsAckFrameType(), "AckFrameType should be recognized as ACK") + require.True(t, FrameTypeAckECN.IsAckFrameType(), "AckECNFrameType should be recognized as ACK") + require.False(t, FrameTypePing.IsAckFrameType(), "PingFrameType should not be recognized as ACK") + require.False(t, FrameType(0x10).IsAckFrameType(), "MaxDataFrameType should not be recognized as ACK") +} + +func TestIsDatagramFrameType(t *testing.T) { + require.True(t, FrameTypeDatagramNoLength.IsDatagramFrameType(), "DatagramNoLengthFrameType should be recognized as DATAGRAM") + require.True(t, FrameTypeDatagramWithLength.IsDatagramFrameType(), "DatagramWithLengthFrameType should be recognized as DATAGRAM") + require.False(t, FrameTypePing.IsDatagramFrameType(), "PingFrameType should not be recognized as DATAGRAM") + require.False(t, FrameType(0x1e).IsDatagramFrameType(), "HandshakeDoneFrameType should not be recognized as DATAGRAM") +} diff --git a/third_party/quic-go/internal/wire/handshake_done_frame.go b/third_party/quic-go/internal/wire/handshake_done_frame.go new file mode 100644 index 0000000..b9bed89 --- /dev/null +++ b/third_party/quic-go/internal/wire/handshake_done_frame.go @@ -0,0 +1,17 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" +) + +// A HandshakeDoneFrame is a HANDSHAKE_DONE frame +type HandshakeDoneFrame struct{} + +func (f *HandshakeDoneFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + return append(b, byte(FrameTypeHandshakeDone)), nil +} + +// Length of a written frame +func (f *HandshakeDoneFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 +} diff --git a/third_party/quic-go/internal/wire/handshake_done_frame_test.go b/third_party/quic-go/internal/wire/handshake_done_frame_test.go new file mode 100644 index 0000000..b30087c --- /dev/null +++ b/third_party/quic-go/internal/wire/handshake_done_frame_test.go @@ -0,0 +1,16 @@ +package wire + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/stretchr/testify/require" +) + +func TestWriteHandshakeDoneSampleFrame(t *testing.T) { + frame := HandshakeDoneFrame{} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, []byte{byte(FrameTypeHandshakeDone)}, b) + require.Equal(t, protocol.ByteCount(1), frame.Length(protocol.Version1)) +} diff --git a/third_party/quic-go/internal/wire/header.go b/third_party/quic-go/internal/wire/header.go new file mode 100644 index 0000000..2e69889 --- /dev/null +++ b/third_party/quic-go/internal/wire/header.go @@ -0,0 +1,302 @@ +package wire + +import ( + "encoding/binary" + "errors" + "fmt" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// ParseConnectionID parses the destination connection ID of a packet. +func ParseConnectionID(data []byte, shortHeaderConnIDLen int) (protocol.ConnectionID, error) { + if len(data) == 0 { + return protocol.ConnectionID{}, io.EOF + } + if !IsLongHeaderPacket(data[0]) { + if len(data) < shortHeaderConnIDLen+1 { + return protocol.ConnectionID{}, io.EOF + } + return protocol.ParseConnectionID(data[1 : 1+shortHeaderConnIDLen]), nil + } + if len(data) < 6 { + return protocol.ConnectionID{}, io.EOF + } + destConnIDLen := int(data[5]) + if destConnIDLen > protocol.MaxConnIDLen { + return protocol.ConnectionID{}, protocol.ErrInvalidConnectionIDLen + } + if len(data) < 6+destConnIDLen { + return protocol.ConnectionID{}, io.EOF + } + return protocol.ParseConnectionID(data[6 : 6+destConnIDLen]), nil +} + +// ParseArbitraryLenConnectionIDs parses the most general form of a Long Header packet, +// using only the version-independent packet format as described in Section 5.1 of RFC 8999: +// https://datatracker.ietf.org/doc/html/rfc8999#section-5.1. +// This function should only be called on Long Header packets for which we don't support the version. +func ParseArbitraryLenConnectionIDs(data []byte) (bytesParsed int, dest, src protocol.ArbitraryLenConnectionID, _ error) { + startLen := len(data) + if len(data) < 6 { + return 0, nil, nil, io.EOF + } + data = data[5:] // skip first byte and version field + destConnIDLen := data[0] + data = data[1:] + destConnID := make(protocol.ArbitraryLenConnectionID, destConnIDLen) + if len(data) < int(destConnIDLen)+1 { + return 0, nil, nil, io.EOF + } + copy(destConnID, data) + data = data[destConnIDLen:] + srcConnIDLen := data[0] + data = data[1:] + if len(data) < int(srcConnIDLen) { + return 0, nil, nil, io.EOF + } + srcConnID := make(protocol.ArbitraryLenConnectionID, srcConnIDLen) + copy(srcConnID, data) + return startLen - len(data) + int(srcConnIDLen), destConnID, srcConnID, nil +} + +func IsPotentialQUICPacket(firstByte byte) bool { + return firstByte&0x40 > 0 +} + +// IsLongHeaderPacket says if this is a Long Header packet +func IsLongHeaderPacket(firstByte byte) bool { + return firstByte&0x80 > 0 +} + +// ParseVersion parses the QUIC version. +// It should only be called for Long Header packets (Short Header packets don't contain a version number). +func ParseVersion(data []byte) (protocol.Version, error) { + if len(data) < 5 { + return 0, io.EOF + } + return protocol.Version(binary.BigEndian.Uint32(data[1:5])), nil +} + +// IsVersionNegotiationPacket says if this is a version negotiation packet +func IsVersionNegotiationPacket(b []byte) bool { + if len(b) < 5 { + return false + } + return IsLongHeaderPacket(b[0]) && b[1] == 0 && b[2] == 0 && b[3] == 0 && b[4] == 0 +} + +// Is0RTTPacket says if this is a 0-RTT packet. +// A packet sent with a version we don't understand can never be a 0-RTT packet. +func Is0RTTPacket(b []byte) bool { + if len(b) < 5 { + return false + } + if !IsLongHeaderPacket(b[0]) { + return false + } + version := protocol.Version(binary.BigEndian.Uint32(b[1:5])) + //nolint:exhaustive // We only need to test QUIC versions that we support. + switch version { + case protocol.Version1: + return b[0]>>4&0b11 == 0b01 + case protocol.Version2: + return b[0]>>4&0b11 == 0b10 + default: + return false + } +} + +var ErrUnsupportedVersion = errors.New("unsupported version") + +// The Header is the version independent part of the header +type Header struct { + typeByte byte + Type protocol.PacketType + + Version protocol.Version + SrcConnectionID protocol.ConnectionID + DestConnectionID protocol.ConnectionID + + Length protocol.ByteCount + + Token []byte + + parsedLen protocol.ByteCount // how many bytes were read while parsing this header +} + +// ParsePacket parses a long header packet. +// The packet is cut according to the length field. +// If we understand the version, the packet is parsed up unto the packet number. +// Otherwise, only the invariant part of the header is parsed. +func ParsePacket(data []byte) (*Header, []byte, []byte, error) { + if len(data) == 0 || !IsLongHeaderPacket(data[0]) { + return nil, nil, nil, errors.New("not a long header packet") + } + hdr, err := parseHeader(data) + if err != nil { + if errors.Is(err, ErrUnsupportedVersion) { + return hdr, nil, nil, err + } + return nil, nil, nil, err + } + if protocol.ByteCount(len(data)) < hdr.ParsedLen()+hdr.Length { + return nil, nil, nil, fmt.Errorf("packet length (%d bytes) is smaller than the expected length (%d bytes)", len(data)-int(hdr.ParsedLen()), hdr.Length) + } + packetLen := int(hdr.ParsedLen() + hdr.Length) + return hdr, data[:packetLen], data[packetLen:], nil +} + +// ParseHeader parses the header: +// * if we understand the version: up to the packet number +// * if not, only the invariant part of the header +func parseHeader(b []byte) (*Header, error) { + if len(b) == 0 { + return nil, io.EOF + } + typeByte := b[0] + + h := &Header{typeByte: typeByte} + l, err := h.parseLongHeader(b[1:]) + h.parsedLen = protocol.ByteCount(l) + 1 + return h, err +} + +func (h *Header) parseLongHeader(b []byte) (int, error) { + startLen := len(b) + if len(b) < 5 { + return 0, io.EOF + } + h.Version = protocol.Version(binary.BigEndian.Uint32(b[:4])) + if h.Version != 0 && h.typeByte&0x40 == 0 { + return startLen - len(b), errors.New("not a QUIC packet") + } + destConnIDLen := int(b[4]) + if destConnIDLen > protocol.MaxConnIDLen { + return startLen - len(b), protocol.ErrInvalidConnectionIDLen + } + b = b[5:] + if len(b) < destConnIDLen+1 { + return startLen - len(b), io.EOF + } + h.DestConnectionID = protocol.ParseConnectionID(b[:destConnIDLen]) + srcConnIDLen := int(b[destConnIDLen]) + if srcConnIDLen > protocol.MaxConnIDLen { + return startLen - len(b), protocol.ErrInvalidConnectionIDLen + } + b = b[destConnIDLen+1:] + if len(b) < srcConnIDLen { + return startLen - len(b), io.EOF + } + h.SrcConnectionID = protocol.ParseConnectionID(b[:srcConnIDLen]) + b = b[srcConnIDLen:] + if h.Version == 0 { // version negotiation packet + return startLen - len(b), nil + } + // If we don't understand the version, we have no idea how to interpret the rest of the bytes + if !protocol.IsSupportedVersion(protocol.SupportedVersions, h.Version) { + return startLen - len(b), ErrUnsupportedVersion + } + + if h.Version == protocol.Version2 { + switch h.typeByte >> 4 & 0b11 { + case 0b00: + h.Type = protocol.PacketTypeRetry + case 0b01: + h.Type = protocol.PacketTypeInitial + case 0b10: + h.Type = protocol.PacketType0RTT + case 0b11: + h.Type = protocol.PacketTypeHandshake + } + } else { + switch h.typeByte >> 4 & 0b11 { + case 0b00: + h.Type = protocol.PacketTypeInitial + case 0b01: + h.Type = protocol.PacketType0RTT + case 0b10: + h.Type = protocol.PacketTypeHandshake + case 0b11: + h.Type = protocol.PacketTypeRetry + } + } + + if h.Type == protocol.PacketTypeRetry { + tokenLen := len(b) - 16 + if tokenLen <= 0 { + return startLen - len(b), io.EOF + } + h.Token = make([]byte, tokenLen) + copy(h.Token, b[:tokenLen]) + return startLen - len(b) + tokenLen + 16, nil + } + + if h.Type == protocol.PacketTypeInitial { + tokenLen, n, err := quicvarint.Parse(b) + if err != nil { + return startLen - len(b), err + } + b = b[n:] + if tokenLen > uint64(len(b)) { + return startLen - len(b), io.EOF + } + h.Token = make([]byte, tokenLen) + copy(h.Token, b[:tokenLen]) + b = b[tokenLen:] + } + + pl, n, err := quicvarint.Parse(b) + if err != nil { + return 0, err + } + h.Length = protocol.ByteCount(pl) + return startLen - len(b) + n, nil +} + +// ParsedLen returns the number of bytes that were consumed when parsing the header +func (h *Header) ParsedLen() protocol.ByteCount { + return h.parsedLen +} + +// ParseExtended parses the version dependent part of the header. +// The Reader has to be set such that it points to the first byte of the header. +func (h *Header) ParseExtended(data []byte) (*ExtendedHeader, error) { + extHdr := h.toExtendedHeader() + reservedBitsValid, err := extHdr.parse(data) + if err != nil { + return nil, err + } + if !reservedBitsValid { + return extHdr, ErrInvalidReservedBits + } + return extHdr, nil +} + +func (h *Header) toExtendedHeader() *ExtendedHeader { + return &ExtendedHeader{Header: *h} +} + +// PacketType is the type of the packet, for logging purposes +func (h *Header) PacketType() string { + return h.Type.String() +} + +func readPacketNumber(data []byte, pnLen protocol.PacketNumberLen) (protocol.PacketNumber, error) { + var pn protocol.PacketNumber + switch pnLen { + case protocol.PacketNumberLen1: + pn = protocol.PacketNumber(data[0]) + case protocol.PacketNumberLen2: + pn = protocol.PacketNumber(binary.BigEndian.Uint16(data[:2])) + case protocol.PacketNumberLen3: + pn = protocol.PacketNumber(uint32(data[2]) + uint32(data[1])<<8 + uint32(data[0])<<16) + case protocol.PacketNumberLen4: + pn = protocol.PacketNumber(binary.BigEndian.Uint32(data[:4])) + default: + return 0, fmt.Errorf("invalid packet number length: %d", pnLen) + } + return pn, nil +} diff --git a/third_party/quic-go/internal/wire/header_test.go b/third_party/quic-go/internal/wire/header_test.go new file mode 100644 index 0000000..43c8e55 --- /dev/null +++ b/third_party/quic-go/internal/wire/header_test.go @@ -0,0 +1,793 @@ +package wire + +import ( + "bytes" + "crypto/rand" + "encoding/binary" + "io" + mrand "math/rand/v2" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" + + "github.com/stretchr/testify/require" +) + +func TestParseConnIDLongHeaderPacket(t *testing.T) { + b, err := (&ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6}), + Version: protocol.Version1, + }, + PacketNumberLen: 2, + }).Append(nil, protocol.Version1) + require.NoError(t, err) + connID, err := ParseConnectionID(b, 8) + require.NoError(t, err) + require.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), connID) +} + +func TestParseConnIDTooLong(t *testing.T) { + b := []byte{0x80, 0, 0, 0, 0} + binary.BigEndian.PutUint32(b[1:], uint32(protocol.Version1)) + b = append(b, 21) // dest conn id len + b = append(b, make([]byte, 21)...) + _, err := ParseConnectionID(b, 4) + require.Error(t, err) + require.ErrorIs(t, err, protocol.ErrInvalidConnectionIDLen) +} + +func TestParseConnIDEOFLongHeader(t *testing.T) { + b, err := (&ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad, 0x13, 0x37}), + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 8, 9}), + Version: protocol.Version1, + }, + PacketNumberLen: 2, + }).Append(nil, protocol.Version1) + require.NoError(t, err) + data := b[:len(b)-2] // cut the packet number + _, err = ParseConnectionID(data, 8) + require.NoError(t, err) + for i := range 1 /* first byte */ + 4 /* version */ + 1 /* conn ID lengths */ + 6 { + b := make([]byte, i) + copy(b, data[:i]) + _, err := ParseConnectionID(b, 8) + require.Error(t, err) + require.ErrorIs(t, err, io.EOF) + } +} + +func TestIs0RTT(t *testing.T) { + t.Run("QUIC v1", func(t *testing.T) { + zeroRTTHeader := make([]byte, 5) + zeroRTTHeader[0] = 0x80 | 0b01<<4 + binary.BigEndian.PutUint32(zeroRTTHeader[1:], uint32(protocol.Version1)) + + require.True(t, Is0RTTPacket(zeroRTTHeader)) + require.False(t, Is0RTTPacket(zeroRTTHeader[:4])) // too short + require.False(t, Is0RTTPacket([]byte{zeroRTTHeader[0], 1, 2, 3, 4})) // unknown version + require.False(t, Is0RTTPacket([]byte{zeroRTTHeader[0] | 0x80, 1, 2, 3, 4})) // short header + require.True(t, Is0RTTPacket(append(zeroRTTHeader, []byte("foobar")...))) + }) + + t.Run("QUIC v2", func(t *testing.T) { + zeroRTTHeader := make([]byte, 5) + zeroRTTHeader[0] = 0x80 | 0b10<<4 + binary.BigEndian.PutUint32(zeroRTTHeader[1:], uint32(protocol.Version2)) + + require.True(t, Is0RTTPacket(zeroRTTHeader)) + require.False(t, Is0RTTPacket(zeroRTTHeader[:4])) // too short + require.False(t, Is0RTTPacket([]byte{zeroRTTHeader[0], 1, 2, 3, 4})) // unknown version + require.False(t, Is0RTTPacket([]byte{zeroRTTHeader[0] | 0x80, 1, 2, 3, 4})) // short header + require.True(t, Is0RTTPacket(append(zeroRTTHeader, []byte("foobar")...))) + }) +} + +func TestParseVersion(t *testing.T) { + b := []byte{0x80, 0xde, 0xad, 0xbe, 0xef} + v, err := ParseVersion(b) + require.NoError(t, err) + require.Equal(t, protocol.Version(0xdeadbeef), v) + + for i := range b { + _, err := ParseVersion(b[:i]) + require.ErrorIs(t, err, io.EOF) + } +} + +func TestParseArbitraryLengthConnectionIDs(t *testing.T) { + generateConnID := func(l int) protocol.ArbitraryLenConnectionID { + c := make(protocol.ArbitraryLenConnectionID, l) + rand.Read(c) + return c + } + + src := generateConnID(mrand.IntN(255) + 1) + dest := generateConnID(mrand.IntN(255) + 1) + b := []byte{0x80, 1, 2, 3, 4} + b = append(b, uint8(dest.Len())) + b = append(b, dest.Bytes()...) + b = append(b, uint8(src.Len())) + b = append(b, src.Bytes()...) + l := len(b) + b = append(b, []byte("foobar")...) // add some payload + + parsed, d, s, err := ParseArbitraryLenConnectionIDs(b) + require.Equal(t, l, parsed) + require.NoError(t, err) + require.Equal(t, src, s) + require.Equal(t, dest, d) + + for i := range b[:l] { + _, _, _, err := ParseArbitraryLenConnectionIDs(b[:i]) + require.ErrorIs(t, err, io.EOF) + } +} + +func TestIdentifyVersionNegotiationPackets(t *testing.T) { + require.True(t, IsVersionNegotiationPacket([]byte{0x80 | 0x56, 0, 0, 0, 0})) + require.False(t, IsVersionNegotiationPacket([]byte{0x56, 0, 0, 0, 0})) + require.False(t, IsVersionNegotiationPacket([]byte{0x80, 1, 0, 0, 0})) + require.False(t, IsVersionNegotiationPacket([]byte{0x80, 0, 1, 0, 0})) + require.False(t, IsVersionNegotiationPacket([]byte{0x80, 0, 0, 1, 0})) + require.False(t, IsVersionNegotiationPacket([]byte{0x80, 0, 0, 0, 1})) +} + +func TestVersionNegotiationPacketEOF(t *testing.T) { + vnp := []byte{0x80, 0, 0, 0, 0} + for i := range vnp { + require.False(t, IsVersionNegotiationPacket(vnp[:i])) + } +} + +func TestParseLongHeader(t *testing.T) { + destConnID := protocol.ParseConnectionID([]byte{9, 8, 7, 6, 5, 4, 3, 2, 1}) + srcConnID := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}) + data := []byte{0xc0 ^ 0x3} + data = appendVersion(data, protocol.Version1) + data = append(data, 0x9) // dest conn id length + data = append(data, destConnID.Bytes()...) + data = append(data, 0x4) // src conn id length + data = append(data, srcConnID.Bytes()...) + data = append(data, encodeVarInt(6)...) // token length + data = append(data, []byte("foobar")...) // token + data = append(data, encodeVarInt(10)...) // length + hdrLen := len(data) + data = append(data, []byte{0, 0, 0xbe, 0xef}...) // packet number + data = append(data, []byte("foobar")...) + require.False(t, IsVersionNegotiationPacket(data)) + + hdr, pdata, rest, err := ParsePacket(data) + require.NoError(t, err) + require.Equal(t, data, pdata) + require.Equal(t, destConnID, hdr.DestConnectionID) + require.Equal(t, srcConnID, hdr.SrcConnectionID) + require.Equal(t, protocol.PacketTypeInitial, hdr.Type) + require.Equal(t, []byte("foobar"), hdr.Token) + require.Equal(t, protocol.ByteCount(10), hdr.Length) + require.Equal(t, protocol.Version1, hdr.Version) + require.Empty(t, rest) + extHdr, err := hdr.ParseExtended(data) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumberLen4, extHdr.PacketNumberLen) + require.Equal(t, protocol.PacketNumber(0xbeef), extHdr.PacketNumber) + require.Equal(t, hdrLen, int(hdr.ParsedLen())) + require.Equal(t, hdr.ParsedLen()+4, extHdr.ParsedLen()) +} + +func TestErrorIfReservedBitNotSet(t *testing.T) { + data := []byte{ + 0x80 | 0x2<<4, + 0x11, // connection ID lengths + 0xde, 0xca, 0xfb, 0xad, // dest conn ID + 0xde, 0xad, 0xbe, 0xef, // src conn ID + } + _, _, _, err := ParsePacket(data) + require.EqualError(t, err, "not a QUIC packet") +} + +func TestStopParsingWhenEncounteringUnsupportedVersion(t *testing.T) { + data := []byte{ + 0xc0, + 0xde, 0xad, 0xbe, 0xef, + 0x8, // dest conn ID len + 0x1, 0x2, 0x3, 0x4, 0x5, 0x6, 0x7, 0x8, // dest conn ID + 0x8, // src conn ID len + 0x8, 0x7, 0x6, 0x5, 0x4, 0x3, 0x2, 0x1, // src conn ID + 'f', 'o', 'o', 'b', 'a', 'r', // unspecified bytes + } + hdr, _, rest, err := ParsePacket(data) + require.EqualError(t, err, ErrUnsupportedVersion.Error()) + require.Equal(t, protocol.Version(0xdeadbeef), hdr.Version) + require.Equal(t, protocol.ParseConnectionID([]byte{0x1, 0x2, 0x3, 0x4, 0x5, 0x6, 0x7, 0x8}), hdr.DestConnectionID) + require.Equal(t, protocol.ParseConnectionID([]byte{0x8, 0x7, 0x6, 0x5, 0x4, 0x3, 0x2, 0x1}), hdr.SrcConnectionID) + require.Empty(t, rest) +} + +func TestParseLongHeaderWithoutDestinationConnectionID(t *testing.T) { + data := []byte{0xc0 ^ 0x1<<4} + data = appendVersion(data, protocol.Version1) + data = append(data, 0) // dest conn ID len + data = append(data, 4) // src conn ID len + data = append(data, []byte{0xde, 0xad, 0xbe, 0xef}...) // source connection ID + data = append(data, encodeVarInt(0)...) // length + data = append(data, []byte{0xde, 0xca, 0xfb, 0xad}...) + hdr, _, _, err := ParsePacket(data) + require.NoError(t, err) + require.Equal(t, protocol.PacketType0RTT, hdr.Type) + require.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), hdr.SrcConnectionID) + require.Zero(t, hdr.DestConnectionID) +} + +func TestParseLongHeaderWithoutSourceConnectionID(t *testing.T) { + data := []byte{0xc0 ^ 0x2<<4} + data = appendVersion(data, protocol.Version1) + data = append(data, 10) // dest conn ID len + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}...) // dest connection ID + data = append(data, 0) // src conn ID len + data = append(data, encodeVarInt(0)...) // length + data = append(data, []byte{0xde, 0xca, 0xfb, 0xad}...) + hdr, _, _, err := ParsePacket(data) + require.NoError(t, err) + require.Zero(t, hdr.SrcConnectionID) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), hdr.DestConnectionID) +} + +func TestErrorOnTooLongDestinationConnectionID(t *testing.T) { + data := []byte{0xc0 ^ 0x2<<4} + data = appendVersion(data, protocol.Version1) + data = append(data, 21) // dest conn ID len + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21}...) // dest connection ID + data = append(data, 0x0) // src conn ID len + data = append(data, encodeVarInt(0)...) // length + data = append(data, []byte{0xde, 0xca, 0xfb, 0xad}...) + _, _, _, err := ParsePacket(data) + require.EqualError(t, err, protocol.ErrInvalidConnectionIDLen.Error()) +} + +func TestParseLongHeaderWith2BytePacketNumber(t *testing.T) { + data := []byte{0xc0 ^ 0x1} + data = appendVersion(data, protocol.Version1) // version number + data = append(data, []byte{0x0, 0x0}...) // connection ID lengths + data = append(data, encodeVarInt(0)...) // token length + data = append(data, encodeVarInt(0)...) // length + data = append(data, []byte{0x1, 0x23}...) + + hdr, _, _, err := ParsePacket(data) + require.NoError(t, err) + extHdr, err := hdr.ParseExtended(data) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumber(0x123), extHdr.PacketNumber) + require.Equal(t, protocol.PacketNumberLen2, extHdr.PacketNumberLen) + require.Equal(t, len(data), int(extHdr.ParsedLen())) +} + +func TestParseRetryPacket(t *testing.T) { + for _, version := range []protocol.Version{protocol.Version1, protocol.Version2} { + t.Run(version.String(), func(t *testing.T) { + var packetType byte + if version == protocol.Version1 { + packetType = 0b11 << 4 + } else { + packetType = 0b00 << 4 + } + data := []byte{0xc0 | packetType | (10 - 3) /* connection ID length */} + data = appendVersion(data, version) + data = append(data, []byte{6}...) // dest conn ID len + data = append(data, []byte{6, 5, 4, 3, 2, 1}...) // dest conn ID + data = append(data, []byte{10}...) // src conn ID len + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}...) // source connection ID + data = append(data, []byte{'f', 'o', 'o', 'b', 'a', 'r'}...) // token + data = append(data, []byte{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}...) + hdr, pdata, rest, err := ParsePacket(data) + require.NoError(t, err) + require.Equal(t, protocol.PacketTypeRetry, hdr.Type) + require.Equal(t, version, hdr.Version) + require.Equal(t, protocol.ParseConnectionID([]byte{6, 5, 4, 3, 2, 1}), hdr.DestConnectionID) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), hdr.SrcConnectionID) + require.Equal(t, []byte("foobar"), hdr.Token) + require.Equal(t, data, pdata) + require.Empty(t, rest) + }) + } +} + +func TestRetryPacketTooShortForIntegrityTag(t *testing.T) { + data := []byte{0xc0 | 0x3<<4 | (10 - 3) /* connection ID length */} + data = appendVersion(data, protocol.Version1) + data = append(data, []byte{0, 0}...) // conn ID lens + data = append(data, []byte{'f', 'o', 'o', 'b', 'a', 'r'}...) // token + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}...) + // this results in a token length of 0 + _, _, _, err := ParsePacket(data) + require.Equal(t, io.EOF, err) +} + +func TestTokenLengthTooLarge(t *testing.T) { + data := []byte{0xc0 ^ 0x1} + data = appendVersion(data, protocol.Version1) + data = append(data, 0x0) // connection ID lengths + data = append(data, encodeVarInt(4)...) // token length: 4 bytes (1 byte too long) + data = append(data, encodeVarInt(0x42)...) // length, 1 byte + data = append(data, []byte{0x12, 0x34}...) // packet number + + _, _, _, err := ParsePacket(data) + require.Equal(t, io.EOF, err) +} + +func TestErrorOn5thOr6thBitSet(t *testing.T) { + data := []byte{0xc0 | 0x2<<4 | 0x8 /* set the 5th bit */ | 0x1 /* 2 byte packet number */} + data = appendVersion(data, protocol.Version1) + data = append(data, []byte{0x0, 0x0}...) // connection ID lengths + data = append(data, encodeVarInt(2)...) // length + data = append(data, []byte{0x12, 0x34}...) // packet number + hdr, _, _, err := ParsePacket(data) + require.NoError(t, err) + require.Equal(t, protocol.PacketTypeHandshake, hdr.Type) + extHdr, err := hdr.ParseExtended(data) + require.EqualError(t, err, ErrInvalidReservedBits.Error()) + require.NotNil(t, extHdr) + require.Equal(t, protocol.PacketNumber(0x1234), extHdr.PacketNumber) +} + +func TestHeaderEOF(t *testing.T) { + data := []byte{0xc0 ^ 0x2<<4} + data = appendVersion(data, protocol.Version1) + data = append(data, 0x8) // dest conn ID len + data = append(data, []byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}...) // dest conn ID + data = append(data, 0x8) // src conn ID len + data = append(data, []byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}...) // src conn ID + for i := 1; i < len(data); i++ { + _, _, _, err := ParsePacket(data[:i]) + require.Equal(t, io.EOF, err) + } +} + +func TestParseExtendedHeaderEOF(t *testing.T) { + data := []byte{0xc0 | 0x2<<4 | 0x3} + data = appendVersion(data, protocol.Version1) + data = append(data, []byte{0x0, 0x0}...) // connection ID lengths + data = append(data, encodeVarInt(0)...) // length + hdrLen := len(data) + data = append(data, []byte{0xde, 0xad, 0xbe, 0xef}...) // packet number + for i := hdrLen; i < len(data); i++ { + b := data[:i] + hdr, _, _, err := ParsePacket(b) + require.NoError(t, err) + _, err = hdr.ParseExtended(b) + require.Equal(t, io.EOF, err) + } +} + +func TestParseRetryEOF(t *testing.T) { + data := []byte{0xc0 ^ 0x3<<4} + data = appendVersion(data, protocol.Version1) + data = append(data, []byte{0x0, 0x0}...) // connection ID lengths + data = append(data, 0xa) // Orig Destination Connection ID length + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}...) // source connection ID + hdrLen := len(data) + for i := hdrLen; i < len(data); i++ { + data = data[:i] + hdr, _, _, err := ParsePacket(data) + require.NoError(t, err) + _, err = hdr.ParseExtended(data) + require.Equal(t, io.EOF, err) + } +} + +func TestCoalescedPacketParsing(t *testing.T) { + hdr := Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + Length: 2 + 6, + Version: protocol.Version1, + } + b, err := (&ExtendedHeader{ + Header: hdr, + PacketNumber: 0x1337, + PacketNumberLen: 2, + }).Append(nil, protocol.Version1) + require.NoError(t, err) + hdrRaw := append([]byte{}, b...) + b = append(b, []byte("foobar")...) // payload of the first packet + b = append(b, []byte("raboof")...) // second packet + parsedHdr, data, rest, err := ParsePacket(b) + require.NoError(t, err) + require.Equal(t, hdr.Type, parsedHdr.Type) + require.Equal(t, hdr.DestConnectionID, parsedHdr.DestConnectionID) + require.Equal(t, append(hdrRaw, []byte("foobar")...), data) + require.Equal(t, []byte("raboof"), rest) +} + +func TestCoalescedPacketErrorOnTooSmallPacketNumber(t *testing.T) { + b, err := (&ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + Length: 3, + Version: protocol.Version1, + }, + PacketNumber: 0x1337, + PacketNumberLen: 2, + }).Append(nil, protocol.Version1) + require.NoError(t, err) + _, _, _, err = ParsePacket(b) + require.Error(t, err) + require.Contains(t, err.Error(), "packet length (2 bytes) is smaller than the expected length (3 bytes)") +} + +func TestCoalescedPacketErrorOnTooSmallPayload(t *testing.T) { + b, err := (&ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + Length: 1000, + Version: protocol.Version1, + }, + PacketNumber: 0x1337, + PacketNumberLen: 2, + }).Append(nil, protocol.Version1) + require.NoError(t, err) + b = append(b, make([]byte, 500-2 /* for packet number length */)...) + _, _, _, err = ParsePacket(b) + require.EqualError(t, err, "packet length (500 bytes) is smaller than the expected length (1000 bytes)") +} + +func TestDistinguishesLongAndShortHeaderPackets(t *testing.T) { + require.False(t, IsLongHeaderPacket(0x40)) + require.True(t, IsLongHeaderPacket(0x80^0x40^0x12)) +} + +func TestPacketTypeForLogging(t *testing.T) { + require.Equal(t, "Initial", (&Header{Type: protocol.PacketTypeInitial}).PacketType()) + require.Equal(t, "Handshake", (&Header{Type: protocol.PacketTypeHandshake}).PacketType()) +} + +func BenchmarkIs0RTTPacket(b *testing.B) { + src := mrand.NewChaCha8([32]byte{'f', 'o', 'o', 'b', 'a', 'r'}) + random := mrand.New(src) + packets := make([][]byte, 1024) + for i := range len(packets) { + packets[i] = make([]byte, random.IntN(256)) + src.Read(packets[i]) + } + + var i int + for b.Loop() { + Is0RTTPacket(packets[i%len(packets)]) + i++ + } +} + +func BenchmarkParseInitial(b *testing.B) { + b.Run("without token", func(b *testing.B) { + benchmarkInitialPacketParsing(b, nil) + }) + b.Run("with token", func(b *testing.B) { + token := make([]byte, 32) + rand.Read(token) + benchmarkInitialPacketParsing(b, token) + }) +} + +func benchmarkInitialPacketParsing(b *testing.B, token []byte) { + b.ReportAllocs() + + hdr := Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + SrcConnectionID: protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), + Length: 1000, + Token: token, + Version: protocol.Version1, + } + data, err := (&ExtendedHeader{ + Header: hdr, + PacketNumber: 0x1337, + PacketNumberLen: 4, + }).Append(nil, protocol.Version1) + if err != nil { + b.Fatal(err) + } + data = append(data, make([]byte, 1000)...) + + for b.Loop() { + h, _, _, err := ParsePacket(data) + if err != nil { + b.Fatal(err) + } + if h.Type != hdr.Type || h.DestConnectionID != hdr.DestConnectionID || h.SrcConnectionID != hdr.SrcConnectionID || + !bytes.Equal(h.Token, hdr.Token) { + b.Fatalf("headers don't match: %v vs %v", h, hdr) + } + } +} + +func BenchmarkParseRetry(b *testing.B) { + b.ReportAllocs() + + token := make([]byte, 64) + rand.Read(token) + hdr := &ExtendedHeader{ + Header: Header{ + Type: protocol.PacketTypeRetry, + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + DestConnectionID: protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), + Token: token, + Version: protocol.Version1, + }, + } + data, err := hdr.Append(nil, hdr.Version) + if err != nil { + b.Fatal(err) + } + + for b.Loop() { + h, _, _, err := ParsePacket(data) + if err != nil { + b.Fatal(err) + } + if h.Type != hdr.Type || h.DestConnectionID != hdr.DestConnectionID || h.SrcConnectionID != hdr.SrcConnectionID || + !bytes.Equal(h.Token, hdr.Token[:len(hdr.Token)-16]) { + b.Fatalf("headers don't match: %#v vs %#v", h, hdr) + } + } +} + +func BenchmarkArbitraryHeaderParsing(b *testing.B) { + b.Run("dest 8/ src 10", func(b *testing.B) { benchmarkArbitraryHeaderParsing(b, 8, 10) }) + b.Run("dest 20 / src 20", func(b *testing.B) { benchmarkArbitraryHeaderParsing(b, 20, 20) }) + b.Run("dest 100 / src 150", func(b *testing.B) { benchmarkArbitraryHeaderParsing(b, 100, 150) }) +} + +func benchmarkArbitraryHeaderParsing(b *testing.B, destLen, srcLen int) { + destConnID := make([]byte, destLen) + rand.Read(destConnID) + srcConnID := make([]byte, srcLen) + rand.Read(srcConnID) + buf := []byte{0x80, 1, 2, 3, 4} + buf = append(buf, uint8(destLen)) + buf = append(buf, destConnID...) + buf = append(buf, uint8(srcLen)) + buf = append(buf, srcConnID...) + + b.ReportAllocs() + for b.Loop() { + parsed, d, s, err := ParseArbitraryLenConnectionIDs(buf) + if err != nil { + b.Fatal(err) + } + if parsed != len(buf) { + b.Fatal("expected to parse entire slice") + } + if !bytes.Equal(destConnID, d.Bytes()) { + b.Fatalf("destination connection IDs don't match: %v vs %v", destConnID, d.Bytes()) + } + if !bytes.Equal(srcConnID, s.Bytes()) { + b.Fatalf("source connection IDs don't match: %v vs %v", srcConnID, s.Bytes()) + } + } +} + +type discardLogger struct{} + +func (discardLogger) SetLogLevel(utils.LogLevel) {} +func (discardLogger) SetLogTimeFormat(string) {} +func (discardLogger) WithPrefix(string) utils.Logger { return discardLogger{} } +func (discardLogger) Debug() bool { return false } +func (discardLogger) Errorf(string, ...any) {} +func (discardLogger) Infof(string, ...any) {} +func (discardLogger) Debugf(string, ...any) {} + +func FuzzHeaderParser(f *testing.F) { + corpus := ossfuzzseeds.New(f) + + addLongHeader := func(hdr *ExtendedHeader) { + b, err := hdr.Append(nil, hdr.Version) + require.NoError(f, err) + if hdr.Type == protocol.PacketTypeRetry { + b = append(b, make([]byte, 16)...) // Retry Integrity Tag + } + if hdr.Length > 0 { + b = append(b, make([]byte, hdr.Length)...) + } + corpus.Add(uint8(hdr.DestConnectionID.Len()), b) + } + + for _, v := range []protocol.Version{protocol.Version1, protocol.Version2} { + // Initial without token + addLongHeader(&ExtendedHeader{ + Header: Header{ + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Type: protocol.PacketTypeInitial, + Length: 10, + Version: v, + }, + PacketNumberLen: protocol.PacketNumberLen2, + PacketNumber: 0x42, + }) + // Initial without token, with zero-length src conn id + addLongHeader(&ExtendedHeader{ + Header: Header{ + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Type: protocol.PacketTypeInitial, + Length: 10, + Version: v, + }, + PacketNumberLen: protocol.PacketNumberLen2, + PacketNumber: 0x42, + }) + // Initial with token + addLongHeader(&ExtendedHeader{ + Header: Header{ + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19}), + Type: protocol.PacketTypeInitial, + Length: 10, + Token: []byte("this is a token"), + Version: v, + }, + PacketNumberLen: protocol.PacketNumberLen4, + PacketNumber: 0xdecafbad, + }) + // Handshake packet + addLongHeader(&ExtendedHeader{ + Header: Header{ + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + Type: protocol.PacketTypeHandshake, + Length: 10, + Version: v, + }, + PacketNumberLen: protocol.PacketNumberLen3, + PacketNumber: 0x1337, + }) + // Handshake packet, with zero-length src conn id + addLongHeader(&ExtendedHeader{ + Header: Header{ + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}), + Type: protocol.PacketTypeHandshake, + Length: 10, + Version: v, + }, + PacketNumberLen: protocol.PacketNumberLen1, + PacketNumber: 0x42, + }) + // 0-RTT packet + addLongHeader(&ExtendedHeader{ + Header: Header{ + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9}), + Type: protocol.PacketType0RTT, + Length: 10, + Version: v, + }, + PacketNumberLen: protocol.PacketNumberLen2, + PacketNumber: 0x42, + }) + // Retry packet + addLongHeader(&ExtendedHeader{ + Header: Header{ + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9}), + Type: protocol.PacketTypeRetry, + Token: []byte("foobar"), + Version: v, + }, + }) + } + // Short header + shortHdr, err := AppendShortHeader(nil, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), 0x1337, protocol.PacketNumberLen2, protocol.KeyPhaseOne) + require.NoError(f, err) + corpus.Add(uint8(8), shortHdr) + // Version Negotiation packets + corpus.Add(uint8(0), ComposeVersionNegotiation( + protocol.ArbitraryLenConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + protocol.ArbitraryLenConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + []protocol.Version{0x1234, 0x5678, 0x9abc, 0xdef0}, + )) + corpus.Add(uint8(0), ComposeVersionNegotiation( + protocol.ArbitraryLenConnectionID([]byte{1, 2, 3}), + protocol.ArbitraryLenConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19}), + []protocol.Version{0xdeadbeef}, + )) + + f.Fuzz(func(t *testing.T, connIDLenRaw uint8, data []byte) { + is0RTTPacket := Is0RTTPacket(data) + + if IsVersionNegotiationPacket(data) { + dest, src, versions, err := ParseVersionNegotiationPacket(data) + if err != nil { + return + } + require.NotEmpty(t, versions, "no versions") + ComposeVersionNegotiation(dest, src, versions) + return + } + + connIDLen := int(connIDLenRaw % 21) + // don't return early on error, we want to fuzz the length checks of other functions as well + connID, _ := ParseConnectionID(data, connIDLen) + + if len(data) == 0 { + return + } + + _ = IsPotentialQUICPacket(data[0]) + + if IsLongHeaderPacket(data[0]) { + ParseVersion(data) + } else { + _, _, _, err := ParsePacket(data) + require.EqualError(t, err, "not a long header packet") + + ParseShortHeader(data, connIDLen) + return + } + + hdr, _, _, err := ParsePacket(data) + if err != nil { + return + } + require.Equal(t, connID, hdr.DestConnectionID, "connection IDs don't match") + if (hdr.Type == protocol.PacketType0RTT) != is0RTTPacket { + t.Fatal("inconsistent 0-RTT packet detection") + } + _ = hdr.PacketType() + + var extHdr *ExtendedHeader + if hdr.Type == protocol.PacketTypeRetry { + extHdr = &ExtendedHeader{Header: *hdr} + } else { + var err error + extHdr, err = hdr.ParseExtended(data) + if err != nil { + return + } + require.Equal(t, hdr.ParsedLen()+protocol.ByteCount(extHdr.PacketNumberLen), extHdr.ParsedLen()) + } + extHdr.Log(discardLogger{}) + // We always use a 2-byte encoding for the Length field in Long Header packets. + // Serializing the header will fail when using a higher value. + if hdr.Length > 16383 { + return + } + b, err := extHdr.Append(nil, hdr.Version) + if err != nil { + // We are able to parse packets with connection IDs longer than 20 bytes, + // but in QUIC version 1 and 2, we don't write headers with longer connection IDs. + if hdr.DestConnectionID.Len() <= protocol.MaxConnIDLen && + hdr.SrcConnectionID.Len() <= protocol.MaxConnIDLen { + t.Fatalf("error writing header: %s", err) + } + return + } + // GetLength is not implemented for Retry packets + if hdr.Type != protocol.PacketTypeRetry { + if expLen := extHdr.GetLength(hdr.Version); expLen != protocol.ByteCount(len(b)) { + t.Fatalf("inconsistent header length: %#v. Expected %d, got %d", extHdr, expLen, len(b)) + } + roundtripHdr, err := parseHeader(b) + require.NoError(t, err) + roundtripExtHdr, err := roundtripHdr.ParseExtended(b) + require.NoError(t, err) + require.Equal(t, extHdr.Type, roundtripExtHdr.Type) + require.Equal(t, extHdr.Version, roundtripExtHdr.Version) + require.Equal(t, extHdr.DestConnectionID, roundtripExtHdr.DestConnectionID) + require.Equal(t, extHdr.SrcConnectionID, roundtripExtHdr.SrcConnectionID) + require.Equal(t, extHdr.Length, roundtripExtHdr.Length) + require.Equal(t, extHdr.Token, roundtripExtHdr.Token) + require.Equal(t, extHdr.PacketNumberLen, roundtripExtHdr.PacketNumberLen) + require.Equal(t, extHdr.PacketNumber, roundtripExtHdr.PacketNumber) + } + }) +} diff --git a/third_party/quic-go/internal/wire/immediate_ack_frame.go b/third_party/quic-go/internal/wire/immediate_ack_frame.go new file mode 100644 index 0000000..1d2a5f8 --- /dev/null +++ b/third_party/quic-go/internal/wire/immediate_ack_frame.go @@ -0,0 +1,18 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// An ImmediateAckFrame is an IMMEDIATE_ACK frame +type ImmediateAckFrame struct{} + +func (f *ImmediateAckFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + return quicvarint.Append(b, uint64(FrameTypeImmediateAck)), nil +} + +// Length of a written frame +func (f *ImmediateAckFrame) Length(_ protocol.Version) protocol.ByteCount { + return protocol.ByteCount(quicvarint.Len(uint64(FrameTypeImmediateAck))) +} diff --git a/third_party/quic-go/internal/wire/immediate_ack_frame_test.go b/third_party/quic-go/internal/wire/immediate_ack_frame_test.go new file mode 100644 index 0000000..0a2484f --- /dev/null +++ b/third_party/quic-go/internal/wire/immediate_ack_frame_test.go @@ -0,0 +1,22 @@ +package wire + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + "github.com/stretchr/testify/require" +) + +func TestImmediateAckFrame(t *testing.T) { + frame := ImmediateAckFrame{} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + + val, l, err := quicvarint.Parse(b) + require.NoError(t, err) + require.Equal(t, uint64(FrameTypeImmediateAck), val) + require.Equal(t, len(b), l) + + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/log.go b/third_party/quic-go/internal/wire/log.go new file mode 100644 index 0000000..4c34fbf --- /dev/null +++ b/third_party/quic-go/internal/wire/log.go @@ -0,0 +1,74 @@ +package wire + +import ( + "fmt" + "strings" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +// LogFrame logs a frame, either sent or received +func LogFrame(logger utils.Logger, frame Frame, sent bool) { + if !logger.Debug() { + return + } + dir := "<-" + if sent { + dir = "->" + } + switch f := frame.(type) { + case *CryptoFrame: + dataLen := protocol.ByteCount(len(f.Data)) + logger.Debugf("\t%s &wire.CryptoFrame{Offset: %d, Data length: %d, Offset + Data length: %d}", dir, f.Offset, dataLen, f.Offset+dataLen) + case *StreamFrame: + logger.Debugf("\t%s &wire.StreamFrame{StreamID: %d, Fin: %t, Offset: %d, Data length: %d, Offset + Data length: %d}", dir, f.StreamID, f.Fin, f.Offset, f.DataLen(), f.Offset+f.DataLen()) + case *ResetStreamFrame: + logger.Debugf("\t%s &wire.ResetStreamFrame{StreamID: %d, ErrorCode: %#x, FinalSize: %d}", dir, f.StreamID, f.ErrorCode, f.FinalSize) + case *AckFrame: + hasECN := f.ECT0 > 0 || f.ECT1 > 0 || f.ECNCE > 0 + var ecn string + if hasECN { + ecn = fmt.Sprintf(", ECT0: %d, ECT1: %d, CE: %d", f.ECT0, f.ECT1, f.ECNCE) + } + if len(f.AckRanges) > 1 { + ackRanges := make([]string, len(f.AckRanges)) + for i, r := range f.AckRanges { + ackRanges[i] = fmt.Sprintf("{Largest: %d, Smallest: %d}", r.Largest, r.Smallest) + } + logger.Debugf("\t%s &wire.AckFrame{LargestAcked: %d, LowestAcked: %d, AckRanges: {%s}, DelayTime: %s%s}", dir, f.LargestAcked(), f.LowestAcked(), strings.Join(ackRanges, ", "), f.DelayTime.String(), ecn) + } else { + logger.Debugf("\t%s &wire.AckFrame{LargestAcked: %d, LowestAcked: %d, DelayTime: %s%s}", dir, f.LargestAcked(), f.LowestAcked(), f.DelayTime.String(), ecn) + } + case *MaxDataFrame: + logger.Debugf("\t%s &wire.MaxDataFrame{MaximumData: %d}", dir, f.MaximumData) + case *MaxStreamDataFrame: + logger.Debugf("\t%s &wire.MaxStreamDataFrame{StreamID: %d, MaximumStreamData: %d}", dir, f.StreamID, f.MaximumStreamData) + case *DataBlockedFrame: + logger.Debugf("\t%s &wire.DataBlockedFrame{MaximumData: %d}", dir, f.MaximumData) + case *StreamDataBlockedFrame: + logger.Debugf("\t%s &wire.StreamDataBlockedFrame{StreamID: %d, MaximumStreamData: %d}", dir, f.StreamID, f.MaximumStreamData) + case *MaxStreamsFrame: + switch f.Type { + case protocol.StreamTypeUni: + logger.Debugf("\t%s &wire.MaxStreamsFrame{Type: uni, MaxStreamNum: %d}", dir, f.MaxStreamNum) + case protocol.StreamTypeBidi: + logger.Debugf("\t%s &wire.MaxStreamsFrame{Type: bidi, MaxStreamNum: %d}", dir, f.MaxStreamNum) + } + case *StreamsBlockedFrame: + switch f.Type { + case protocol.StreamTypeUni: + logger.Debugf("\t%s &wire.StreamsBlockedFrame{Type: uni, MaxStreams: %d}", dir, f.StreamLimit) + case protocol.StreamTypeBidi: + logger.Debugf("\t%s &wire.StreamsBlockedFrame{Type: bidi, MaxStreams: %d}", dir, f.StreamLimit) + } + case *NewConnectionIDFrame: + logger.Debugf("\t%s &wire.NewConnectionIDFrame{SequenceNumber: %d, RetirePriorTo: %d, ConnectionID: %s, StatelessResetToken: %#x}", dir, f.SequenceNumber, f.RetirePriorTo, f.ConnectionID, f.StatelessResetToken) + case *RetireConnectionIDFrame: + logger.Debugf("\t%s &wire.RetireConnectionIDFrame{SequenceNumber: %d}", dir, f.SequenceNumber) + case *NewTokenFrame: + logger.Debugf("\t%s &wire.NewTokenFrame{Token: %#x}", dir, f.Token) + default: + logger.Debugf("\t%s %#v", dir, frame) + } +} diff --git a/third_party/quic-go/internal/wire/log_test.go b/third_party/quic-go/internal/wire/log_test.go new file mode 100644 index 0000000..2b13ed4 --- /dev/null +++ b/third_party/quic-go/internal/wire/log_test.go @@ -0,0 +1,188 @@ +package wire + +import ( + "bytes" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + + "github.com/stretchr/testify/require" +) + +func TestLogFrameNoDebug(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + logger.SetLogLevel(utils.LogLevelInfo) + LogFrame(logger, &ResetStreamFrame{}, true) + require.Zero(t, buf.Len()) +} + +func TestLogSentFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + LogFrame(logger, &ResetStreamFrame{}, true) + require.Contains(t, buf.String(), "\t-> &wire.ResetStreamFrame{StreamID: 0, ErrorCode: 0x0, FinalSize: 0}\n") +} + +func TestLogReceivedFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + LogFrame(logger, &ResetStreamFrame{}, false) + require.Contains(t, buf.String(), "\t<- &wire.ResetStreamFrame{StreamID: 0, ErrorCode: 0x0, FinalSize: 0}\n") +} + +func TestLogCryptoFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &CryptoFrame{ + Offset: 42, + Data: make([]byte, 123), + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.CryptoFrame{Offset: 42, Data length: 123, Offset + Data length: 165}\n") +} + +func TestLogStreamFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &StreamFrame{ + StreamID: 42, + Offset: 1337, + Data: bytes.Repeat([]byte{'f'}, 100), + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.StreamFrame{StreamID: 42, Fin: false, Offset: 1337, Data length: 100, Offset + Data length: 1437}\n") +} + +func TestLogAckFrameWithoutMissingPackets(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &AckFrame{ + AckRanges: []AckRange{{Smallest: 42, Largest: 1337}}, + DelayTime: 1 * time.Millisecond, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.AckFrame{LargestAcked: 1337, LowestAcked: 42, DelayTime: 1ms}\n") +} + +func TestLogAckFrameWithECN(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &AckFrame{ + AckRanges: []AckRange{{Smallest: 42, Largest: 1337}}, + DelayTime: 1 * time.Millisecond, + ECT0: 5, + ECT1: 66, + ECNCE: 777, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.AckFrame{LargestAcked: 1337, LowestAcked: 42, DelayTime: 1ms, ECT0: 5, ECT1: 66, CE: 777}\n") +} + +func TestLogAckFrameWithMissingPackets(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &AckFrame{ + AckRanges: []AckRange{ + {Smallest: 5, Largest: 8}, + {Smallest: 2, Largest: 3}, + }, + DelayTime: 12 * time.Millisecond, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.AckFrame{LargestAcked: 8, LowestAcked: 2, AckRanges: {{Largest: 8, Smallest: 5}, {Largest: 3, Smallest: 2}}, DelayTime: 12ms}\n") +} + +func TestLogMaxStreamsFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &MaxStreamsFrame{ + Type: protocol.StreamTypeBidi, + MaxStreamNum: 42, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.MaxStreamsFrame{Type: bidi, MaxStreamNum: 42}\n") +} + +func TestLogMaxDataFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &MaxDataFrame{ + MaximumData: 42, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.MaxDataFrame{MaximumData: 42}\n") +} + +func TestLogMaxStreamDataFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &MaxStreamDataFrame{ + StreamID: 10, + MaximumStreamData: 42, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.MaxStreamDataFrame{StreamID: 10, MaximumStreamData: 42}\n") +} + +func TestLogDataBlockedFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &DataBlockedFrame{ + MaximumData: 1000, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.DataBlockedFrame{MaximumData: 1000}\n") +} + +func TestLogStreamDataBlockedFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &StreamDataBlockedFrame{ + StreamID: 42, + MaximumStreamData: 1000, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.StreamDataBlockedFrame{StreamID: 42, MaximumStreamData: 1000}\n") +} + +func TestLogStreamsBlockedFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + frame := &StreamsBlockedFrame{ + Type: protocol.StreamTypeBidi, + StreamLimit: 42, + } + LogFrame(logger, frame, false) + require.Contains(t, buf.String(), "\t<- &wire.StreamsBlockedFrame{Type: bidi, MaxStreams: 42}\n") +} + +func TestLogNewConnectionIDFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + LogFrame(logger, &NewConnectionIDFrame{ + SequenceNumber: 42, + RetirePriorTo: 24, + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{0x1, 0x2, 0x3, 0x4, 0x5, 0x6, 0x7, 0x8, 0x9, 0xa, 0xb, 0xc, 0xd, 0xe, 0xf, 0x10}, + }, false) + require.Contains(t, buf.String(), "\t<- &wire.NewConnectionIDFrame{SequenceNumber: 42, RetirePriorTo: 24, ConnectionID: deadbeef, StatelessResetToken: 0x0102030405060708090a0b0c0d0e0f10}") +} + +func TestLogRetireConnectionIDFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + LogFrame(logger, &RetireConnectionIDFrame{SequenceNumber: 42}, false) + require.Contains(t, buf.String(), "\t<- &wire.RetireConnectionIDFrame{SequenceNumber: 42}") +} + +func TestLogNewTokenFrame(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + LogFrame(logger, &NewTokenFrame{ + Token: []byte{0xde, 0xad, 0xbe, 0xef}, + }, true) + require.Contains(t, buf.String(), "\t-> &wire.NewTokenFrame{Token: 0xdeadbeef") +} diff --git a/third_party/quic-go/internal/wire/max_data_frame.go b/third_party/quic-go/internal/wire/max_data_frame.go new file mode 100644 index 0000000..c4f36d0 --- /dev/null +++ b/third_party/quic-go/internal/wire/max_data_frame.go @@ -0,0 +1,33 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A MaxDataFrame carries flow control information for the connection +type MaxDataFrame struct { + MaximumData protocol.ByteCount +} + +// parseMaxDataFrame parses a MAX_DATA frame +func parseMaxDataFrame(b []byte, _ protocol.Version) (*MaxDataFrame, int, error) { + frame := &MaxDataFrame{} + byteOffset, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + frame.MaximumData = protocol.ByteCount(byteOffset) + return frame, l, nil +} + +func (f *MaxDataFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeMaxData)) + b = quicvarint.Append(b, uint64(f.MaximumData)) + return b, nil +} + +// Length of a written frame +func (f *MaxDataFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.MaximumData))) +} diff --git a/third_party/quic-go/internal/wire/max_data_frame_test.go b/third_party/quic-go/internal/wire/max_data_frame_test.go new file mode 100644 index 0000000..d95c456 --- /dev/null +++ b/third_party/quic-go/internal/wire/max_data_frame_test.go @@ -0,0 +1,39 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseMaxDataFrame(t *testing.T) { + data := encodeVarInt(0xdecafbad123456) // byte offset + frame, l, err := parseMaxDataFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.ByteCount(0xdecafbad123456), frame.MaximumData) + require.Equal(t, len(data), l) +} + +func TestParseMaxDataErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdecafbad1234567) // byte offset + _, l, err := parseMaxDataFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseMaxDataFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteMaxDataFrame(t *testing.T) { + f := &MaxDataFrame{MaximumData: 0xdeadbeefcafe} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeMaxData)} + expected = append(expected, encodeVarInt(0xdeadbeefcafe)...) + require.Equal(t, expected, b) + require.Len(t, b, int(f.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/max_stream_data_frame.go b/third_party/quic-go/internal/wire/max_stream_data_frame.go new file mode 100644 index 0000000..231a715 --- /dev/null +++ b/third_party/quic-go/internal/wire/max_stream_data_frame.go @@ -0,0 +1,43 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A MaxStreamDataFrame is a MAX_STREAM_DATA frame +type MaxStreamDataFrame struct { + StreamID protocol.StreamID + MaximumStreamData protocol.ByteCount +} + +func parseMaxStreamDataFrame(b []byte, _ protocol.Version) (*MaxStreamDataFrame, int, error) { + startLen := len(b) + sid, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + offset, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + + return &MaxStreamDataFrame{ + StreamID: protocol.StreamID(sid), + MaximumStreamData: protocol.ByteCount(offset), + }, startLen - len(b), nil +} + +func (f *MaxStreamDataFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeMaxStreamData)) + b = quicvarint.Append(b, uint64(f.StreamID)) + b = quicvarint.Append(b, uint64(f.MaximumStreamData)) + return b, nil +} + +// Length of a written frame +func (f *MaxStreamDataFrame) Length(protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.StreamID))+quicvarint.Len(uint64(f.MaximumStreamData))) +} diff --git a/third_party/quic-go/internal/wire/max_stream_data_frame_test.go b/third_party/quic-go/internal/wire/max_stream_data_frame_test.go new file mode 100644 index 0000000..494b40d --- /dev/null +++ b/third_party/quic-go/internal/wire/max_stream_data_frame_test.go @@ -0,0 +1,46 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseMaxStreamFrame(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // Stream ID + data = append(data, encodeVarInt(0x12345678)...) // Offset + frame, l, err := parseMaxStreamDataFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0xdeadbeef), frame.StreamID) + require.Equal(t, protocol.ByteCount(0x12345678), frame.MaximumStreamData) + require.Equal(t, len(data), l) +} + +func TestParseMaxStreamDataErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // Stream ID + data = append(data, encodeVarInt(0x12345678)...) // Offset + _, l, err := parseMaxStreamDataFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseMaxStreamDataFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteMaxStreamDataFrame(t *testing.T) { + f := &MaxStreamDataFrame{ + StreamID: 0xdecafbad, + MaximumStreamData: 0xdeadbeefcafe42, + } + expected := []byte{byte(FrameTypeMaxStreamData)} + expected = append(expected, encodeVarInt(0xdecafbad)...) + expected = append(expected, encodeVarInt(0xdeadbeefcafe42)...) + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, expected, b) + require.Equal(t, len(b), int(f.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/max_streams_frame.go b/third_party/quic-go/internal/wire/max_streams_frame.go new file mode 100644 index 0000000..12a5f95 --- /dev/null +++ b/third_party/quic-go/internal/wire/max_streams_frame.go @@ -0,0 +1,50 @@ +package wire + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A MaxStreamsFrame is a MAX_STREAMS frame +type MaxStreamsFrame struct { + Type protocol.StreamType + MaxStreamNum protocol.StreamNum +} + +func parseMaxStreamsFrame(b []byte, typ FrameType, _ protocol.Version) (*MaxStreamsFrame, int, error) { + f := &MaxStreamsFrame{} + //nolint:exhaustive // Function will only be called with BidiMaxStreamsFrameType or UniMaxStreamsFrameType + switch typ { + case FrameTypeBidiMaxStreams: + f.Type = protocol.StreamTypeBidi + case FrameTypeUniMaxStreams: + f.Type = protocol.StreamTypeUni + } + streamID, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + f.MaxStreamNum = protocol.StreamNum(streamID) + if f.MaxStreamNum > protocol.MaxStreamCount { + return nil, 0, fmt.Errorf("%d exceeds the maximum stream count", f.MaxStreamNum) + } + return f, l, nil +} + +func (f *MaxStreamsFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + switch f.Type { + case protocol.StreamTypeBidi: + b = append(b, byte(FrameTypeBidiMaxStreams)) + case protocol.StreamTypeUni: + b = append(b, byte(FrameTypeUniMaxStreams)) + } + b = quicvarint.Append(b, uint64(f.MaxStreamNum)) + return b, nil +} + +// Length of a written frame +func (f *MaxStreamsFrame) Length(protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.MaxStreamNum))) +} diff --git a/third_party/quic-go/internal/wire/max_streams_frame_test.go b/third_party/quic-go/internal/wire/max_streams_frame_test.go new file mode 100644 index 0000000..1ae215e --- /dev/null +++ b/third_party/quic-go/internal/wire/max_streams_frame_test.go @@ -0,0 +1,117 @@ +package wire + +import ( + "fmt" + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +func TestParseMaxStreamsFrameBidirectional(t *testing.T) { + data := encodeVarInt(0xdecaf) + f, l, err := parseMaxStreamsFrame(data, FrameTypeBidiMaxStreams, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamTypeBidi, f.Type) + require.EqualValues(t, 0xdecaf, f.MaxStreamNum) + require.Equal(t, len(data), l) +} + +func TestParseMaxStreamsFrameUnidirectional(t *testing.T) { + data := encodeVarInt(0xdecaf) + f, l, err := parseMaxStreamsFrame(data, FrameTypeUniMaxStreams, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamTypeUni, f.Type) + require.EqualValues(t, 0xdecaf, f.MaxStreamNum) + require.Equal(t, len(data), l) +} + +func TestParseMaxStreamsErrorsOnEOF(t *testing.T) { + const typ = 0x1d + data := encodeVarInt(0xdeadbeefcafe13) + _, l, err := parseMaxStreamsFrame(data, typ, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseMaxStreamsFrame(data[:i], typ, protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestParseMaxStreamsMaxValue(t *testing.T) { + for _, streamType := range []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi} { + var streamTypeStr string + if streamType == protocol.StreamTypeUni { + streamTypeStr = "unidirectional" + } else { + streamTypeStr = "bidirectional" + } + t.Run(streamTypeStr, func(t *testing.T) { + f := &MaxStreamsFrame{ + Type: streamType, + MaxStreamNum: protocol.MaxStreamCount, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + typ, l, err := quicvarint.Parse(b) + require.NoError(t, err) + b = b[l:] + frame, _, err := parseMaxStreamsFrame(b, FrameType(typ), protocol.Version1) + require.NoError(t, err) + require.Equal(t, f, frame) + }) + } +} + +func TestParseMaxStreamsErrorsOnTooLargeStreamCount(t *testing.T) { + for _, streamType := range []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi} { + var streamTypeStr string + if streamType == protocol.StreamTypeUni { + streamTypeStr = "unidirectional" + } else { + streamTypeStr = "bidirectional" + } + t.Run(streamTypeStr, func(t *testing.T) { + f := &MaxStreamsFrame{ + Type: streamType, + MaxStreamNum: protocol.MaxStreamCount + 1, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + typ, l, err := quicvarint.Parse(b) + require.NoError(t, err) + b = b[l:] + _, _, err = parseMaxStreamsFrame(b, FrameType(typ), protocol.Version1) + require.EqualError(t, err, fmt.Sprintf("%d exceeds the maximum stream count", protocol.MaxStreamCount+1)) + }) + } +} + +func TestWriteMaxStreamsBidirectional(t *testing.T) { + f := &MaxStreamsFrame{ + Type: protocol.StreamTypeBidi, + MaxStreamNum: 0xdeadbeef, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeBidiMaxStreams)} + expected = append(expected, encodeVarInt(0xdeadbeef)...) + require.Equal(t, expected, b) + require.Len(t, b, int(f.Length(protocol.Version1))) +} + +func TestWriteMaxStreamsUnidirectional(t *testing.T) { + f := &MaxStreamsFrame{ + Type: protocol.StreamTypeUni, + MaxStreamNum: 0xdecafbad, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeUniMaxStreams)} + expected = append(expected, encodeVarInt(0xdecafbad)...) + require.Equal(t, expected, b) + require.Len(t, b, int(f.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/new_connection_id_frame.go b/third_party/quic-go/internal/wire/new_connection_id_frame.go new file mode 100644 index 0000000..58cdb47 --- /dev/null +++ b/third_party/quic-go/internal/wire/new_connection_id_frame.go @@ -0,0 +1,80 @@ +package wire + +import ( + "errors" + "fmt" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A NewConnectionIDFrame is a NEW_CONNECTION_ID frame +type NewConnectionIDFrame struct { + SequenceNumber uint64 + RetirePriorTo uint64 + ConnectionID protocol.ConnectionID + StatelessResetToken protocol.StatelessResetToken +} + +func parseNewConnectionIDFrame(b []byte, _ protocol.Version) (*NewConnectionIDFrame, int, error) { + startLen := len(b) + seq, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + ret, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + if ret > seq { + //nolint:staticcheck // SA1021: Retire Prior To is the name of the field + return nil, 0, fmt.Errorf("Retire Prior To value (%d) larger than Sequence Number (%d)", ret, seq) + } + if len(b) == 0 { + return nil, 0, io.EOF + } + connIDLen := int(b[0]) + b = b[1:] + if connIDLen == 0 { + return nil, 0, errors.New("invalid zero-length connection ID") + } + if connIDLen > protocol.MaxConnIDLen { + return nil, 0, protocol.ErrInvalidConnectionIDLen + } + if len(b) < connIDLen { + return nil, 0, io.EOF + } + frame := &NewConnectionIDFrame{ + SequenceNumber: seq, + RetirePriorTo: ret, + ConnectionID: protocol.ParseConnectionID(b[:connIDLen]), + } + b = b[connIDLen:] + if len(b) < len(frame.StatelessResetToken) { + return nil, 0, io.EOF + } + copy(frame.StatelessResetToken[:], b) + return frame, startLen - len(b) + len(frame.StatelessResetToken), nil +} + +func (f *NewConnectionIDFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeNewConnectionID)) + b = quicvarint.Append(b, f.SequenceNumber) + b = quicvarint.Append(b, f.RetirePriorTo) + connIDLen := f.ConnectionID.Len() + if connIDLen > protocol.MaxConnIDLen { + return nil, fmt.Errorf("invalid connection ID length: %d", connIDLen) + } + b = append(b, uint8(connIDLen)) + b = append(b, f.ConnectionID.Bytes()...) + b = append(b, f.StatelessResetToken[:]...) + return b, nil +} + +// Length of a written frame +func (f *NewConnectionIDFrame) Length(protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(f.SequenceNumber)+quicvarint.Len(f.RetirePriorTo)+1 /* connection ID length */ +f.ConnectionID.Len()) + 16 +} diff --git a/third_party/quic-go/internal/wire/new_connection_id_frame_test.go b/third_party/quic-go/internal/wire/new_connection_id_frame_test.go new file mode 100644 index 0000000..ae8476a --- /dev/null +++ b/third_party/quic-go/internal/wire/new_connection_id_frame_test.go @@ -0,0 +1,88 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseNewConnectionIDFrame(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // sequence number + data = append(data, encodeVarInt(0xcafe)...) // retire prior to + data = append(data, 10) // connection ID length + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}...) // connection ID + data = append(data, []byte("deadbeefdecafbad")...) // stateless reset token + frame, l, err := parseNewConnectionIDFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, uint64(0xdeadbeef), frame.SequenceNumber) + require.Equal(t, uint64(0xcafe), frame.RetirePriorTo) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), frame.ConnectionID) + require.Equal(t, "deadbeefdecafbad", string(frame.StatelessResetToken[:])) + require.Equal(t, len(data), l) +} + +func TestParseNewConnectionIDRetirePriorToLargerThanSequenceNumber(t *testing.T) { + data := encodeVarInt(1000) // sequence number + data = append(data, encodeVarInt(1001)...) // retire prior to + data = append(data, 3) + data = append(data, []byte{1, 2, 3}...) + data = append(data, []byte("deadbeefdecafbad")...) // stateless reset token + _, _, err := parseNewConnectionIDFrame(data, protocol.Version1) + require.EqualError(t, err, "Retire Prior To value (1001) larger than Sequence Number (1000)") +} + +func TestParseNewConnectionIDZeroLengthConnID(t *testing.T) { + data := encodeVarInt(42) // sequence number + data = append(data, encodeVarInt(12)...) // retire prior to + data = append(data, 0) // connection ID length + _, _, err := parseNewConnectionIDFrame(data, protocol.Version1) + require.EqualError(t, err, "invalid zero-length connection ID") +} + +func TestParseNewConnectionIDInvalidConnIDLength(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // sequence number + data = append(data, encodeVarInt(0xcafe)...) // retire prior to + data = append(data, 21) // connection ID length + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21}...) // connection ID + data = append(data, []byte("deadbeefdecafbad")...) // stateless reset token + _, _, err := parseNewConnectionIDFrame(data, protocol.Version1) + require.Equal(t, protocol.ErrInvalidConnectionIDLen, err) +} + +func TestParseNewConnectionIDErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // sequence number + data = append(data, encodeVarInt(0xcafe1234)...) // retire prior to + data = append(data, 10) // connection ID length + data = append(data, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}...) // connection ID + data = append(data, []byte("deadbeefdecafbad")...) // stateless reset token + _, l, err := parseNewConnectionIDFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseNewConnectionIDFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteNewConnectionIDFrame(t *testing.T) { + token := protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15} + frame := &NewConnectionIDFrame{ + SequenceNumber: 0x1337, + RetirePriorTo: 0x42, + ConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6}), + StatelessResetToken: token, + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeNewConnectionID)} + expected = append(expected, encodeVarInt(0x1337)...) + expected = append(expected, encodeVarInt(0x42)...) + expected = append(expected, 6) + expected = append(expected, []byte{1, 2, 3, 4, 5, 6}...) + expected = append(expected, token[:]...) + require.Equal(t, expected, b) + require.Equal(t, int(frame.Length(protocol.Version1)), len(b)) +} diff --git a/third_party/quic-go/internal/wire/new_token_frame.go b/third_party/quic-go/internal/wire/new_token_frame.go new file mode 100644 index 0000000..aa3fbb5 --- /dev/null +++ b/third_party/quic-go/internal/wire/new_token_frame.go @@ -0,0 +1,43 @@ +package wire + +import ( + "errors" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A NewTokenFrame is a NEW_TOKEN frame +type NewTokenFrame struct { + Token []byte +} + +func parseNewTokenFrame(b []byte, _ protocol.Version) (*NewTokenFrame, int, error) { + tokenLen, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + if tokenLen == 0 { + return nil, 0, errors.New("token must not be empty") + } + if uint64(len(b)) < tokenLen { + return nil, 0, io.EOF + } + token := make([]byte, int(tokenLen)) + copy(token, b) + return &NewTokenFrame{Token: token}, l + int(tokenLen), nil +} + +func (f *NewTokenFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeNewToken)) + b = quicvarint.Append(b, uint64(len(f.Token))) + b = append(b, f.Token...) + return b, nil +} + +// Length of a written frame +func (f *NewTokenFrame) Length(protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(len(f.Token)))+len(f.Token)) +} diff --git a/third_party/quic-go/internal/wire/new_token_frame_test.go b/third_party/quic-go/internal/wire/new_token_frame_test.go new file mode 100644 index 0000000..d9bcb6e --- /dev/null +++ b/third_party/quic-go/internal/wire/new_token_frame_test.go @@ -0,0 +1,51 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseNewTokenFrame(t *testing.T) { + token := "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua." + data := encodeVarInt(uint64(len(token))) + data = append(data, token...) + f, l, err := parseNewTokenFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, token, string(f.Token)) + require.Equal(t, len(data), l) +} + +func TestParseNewTokenFrameRejectsEmptyTokens(t *testing.T) { + data := encodeVarInt(0) + _, _, err := parseNewTokenFrame(data, protocol.Version1) + require.EqualError(t, err, "token must not be empty") +} + +func TestParseNewTokenFrameErrorsOnEOFs(t *testing.T) { + token := "Lorem ipsum dolor sit amet, consectetur adipiscing elit" + data := encodeVarInt(uint64(len(token))) + data = append(data, token...) + _, l, err := parseNewTokenFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseNewTokenFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteNewTokenFrame(t *testing.T) { + token := "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat." + f := &NewTokenFrame{Token: []byte(token)} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeNewToken)} + expected = append(expected, encodeVarInt(uint64(len(token)))...) + expected = append(expected, token...) + require.Equal(t, expected, b) + require.Equal(t, len(b), int(f.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/path_challenge_frame.go b/third_party/quic-go/internal/wire/path_challenge_frame.go new file mode 100644 index 0000000..7aa01db --- /dev/null +++ b/third_party/quic-go/internal/wire/path_challenge_frame.go @@ -0,0 +1,32 @@ +package wire + +import ( + "io" + + "github.com/apernet/quic-go/internal/protocol" +) + +// A PathChallengeFrame is a PATH_CHALLENGE frame +type PathChallengeFrame struct { + Data [8]byte +} + +func parsePathChallengeFrame(b []byte, _ protocol.Version) (*PathChallengeFrame, int, error) { + f := &PathChallengeFrame{} + if len(b) < 8 { + return nil, 0, io.EOF + } + copy(f.Data[:], b) + return f, 8, nil +} + +func (f *PathChallengeFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypePathChallenge)) + b = append(b, f.Data[:]...) + return b, nil +} + +// Length of a written frame +func (f *PathChallengeFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 + 8 +} diff --git a/third_party/quic-go/internal/wire/path_challenge_frame_test.go b/third_party/quic-go/internal/wire/path_challenge_frame_test.go new file mode 100644 index 0000000..409db3e --- /dev/null +++ b/third_party/quic-go/internal/wire/path_challenge_frame_test.go @@ -0,0 +1,37 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParsePathChallenge(t *testing.T) { + b := []byte{1, 2, 3, 4, 5, 6, 7, 8} + f, l, err := parsePathChallengeFrame(b, protocol.Version1) + require.NoError(t, err) + require.Equal(t, [8]byte{1, 2, 3, 4, 5, 6, 7, 8}, f.Data) + require.Equal(t, len(b), l) +} + +func TestParsePathChallengeErrorsOnEOFs(t *testing.T) { + data := []byte{1, 2, 3, 4, 5, 6, 7, 8} + _, l, err := parsePathChallengeFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parsePathChallengeFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWritePathChallenge(t *testing.T) { + frame := PathChallengeFrame{Data: [8]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, []byte{byte(FrameTypePathChallenge), 0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/path_response_frame.go b/third_party/quic-go/internal/wire/path_response_frame.go new file mode 100644 index 0000000..4447fcf --- /dev/null +++ b/third_party/quic-go/internal/wire/path_response_frame.go @@ -0,0 +1,32 @@ +package wire + +import ( + "io" + + "github.com/apernet/quic-go/internal/protocol" +) + +// A PathResponseFrame is a PATH_RESPONSE frame +type PathResponseFrame struct { + Data [8]byte +} + +func parsePathResponseFrame(b []byte, _ protocol.Version) (*PathResponseFrame, int, error) { + f := &PathResponseFrame{} + if len(b) < 8 { + return nil, 0, io.EOF + } + copy(f.Data[:], b) + return f, 8, nil +} + +func (f *PathResponseFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypePathResponse)) + b = append(b, f.Data[:]...) + return b, nil +} + +// Length of a written frame +func (f *PathResponseFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 + 8 +} diff --git a/third_party/quic-go/internal/wire/path_response_frame_test.go b/third_party/quic-go/internal/wire/path_response_frame_test.go new file mode 100644 index 0000000..5e04177 --- /dev/null +++ b/third_party/quic-go/internal/wire/path_response_frame_test.go @@ -0,0 +1,37 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParsePathResponse(t *testing.T) { + b := []byte{1, 2, 3, 4, 5, 6, 7, 8} + f, l, err := parsePathResponseFrame(b, protocol.Version1) + require.NoError(t, err) + require.Equal(t, [8]byte{1, 2, 3, 4, 5, 6, 7, 8}, f.Data) + require.Equal(t, len(b), l) +} + +func TestParsePathResponseErrorsOnEOFs(t *testing.T) { + data := []byte{1, 2, 3, 4, 5, 6, 7, 8} + _, l, err := parsePathResponseFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parsePathResponseFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWritePathResponse(t *testing.T) { + frame := PathResponseFrame{Data: [8]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, []byte{byte(FrameTypePathResponse), 0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/ping_frame.go b/third_party/quic-go/internal/wire/ping_frame.go new file mode 100644 index 0000000..b7a82f2 --- /dev/null +++ b/third_party/quic-go/internal/wire/ping_frame.go @@ -0,0 +1,17 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" +) + +// A PingFrame is a PING frame +type PingFrame struct{} + +func (f *PingFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + return append(b, byte(FrameTypePing)), nil +} + +// Length of a written frame +func (f *PingFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 +} diff --git a/third_party/quic-go/internal/wire/ping_frame_test.go b/third_party/quic-go/internal/wire/ping_frame_test.go new file mode 100644 index 0000000..80ea5ab --- /dev/null +++ b/third_party/quic-go/internal/wire/ping_frame_test.go @@ -0,0 +1,17 @@ +package wire + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestWritePingFrame(t *testing.T) { + frame := PingFrame{} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, []byte{0x1}, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/pool.go b/third_party/quic-go/internal/wire/pool.go new file mode 100644 index 0000000..5f59c8e --- /dev/null +++ b/third_party/quic-go/internal/wire/pool.go @@ -0,0 +1,33 @@ +package wire + +import ( + "sync" + + "github.com/apernet/quic-go/internal/protocol" +) + +var pool sync.Pool + +func init() { + pool.New = func() any { + return &StreamFrame{ + Data: make([]byte, 0, protocol.MaxPacketBufferSize), + fromPool: true, + } + } +} + +func GetStreamFrame() *StreamFrame { + f := pool.Get().(*StreamFrame) + return f +} + +func putStreamFrame(f *StreamFrame) { + if !f.fromPool { + return + } + if protocol.ByteCount(cap(f.Data)) != protocol.MaxPacketBufferSize { + panic("wire.PutStreamFrame called with packet of wrong size!") + } + pool.Put(f) +} diff --git a/third_party/quic-go/internal/wire/pool_test.go b/third_party/quic-go/internal/wire/pool_test.go new file mode 100644 index 0000000..d67b463 --- /dev/null +++ b/third_party/quic-go/internal/wire/pool_test.go @@ -0,0 +1,24 @@ +package wire + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGetAndPutStreamFrames(t *testing.T) { + f := GetStreamFrame() + putStreamFrame(f) +} + +func TestPanicOnPuttingStreamFrameWithWrongCapacity(t *testing.T) { + f := GetStreamFrame() + f.Data = []byte("foobar") + require.Panics(t, func() { putStreamFrame(f) }) +} + +func TestAcceptStreamFramesNotFromBuffer(t *testing.T) { + f := &StreamFrame{Data: []byte("foobar")} + putStreamFrame(f) + // No assertion needed as we're just checking it doesn't panic +} diff --git a/third_party/quic-go/internal/wire/reset_stream_frame.go b/third_party/quic-go/internal/wire/reset_stream_frame.go new file mode 100644 index 0000000..84ae818 --- /dev/null +++ b/third_party/quic-go/internal/wire/reset_stream_frame.go @@ -0,0 +1,79 @@ +package wire + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/quicvarint" +) + +// A ResetStreamFrame is a RESET_STREAM or RESET_STREAM_AT frame in QUIC +type ResetStreamFrame struct { + StreamID protocol.StreamID + ErrorCode qerr.StreamErrorCode + FinalSize protocol.ByteCount + ReliableSize protocol.ByteCount +} + +func parseResetStreamFrame(b []byte, isResetStreamAt bool, _ protocol.Version) (*ResetStreamFrame, int, error) { + startLen := len(b) + streamID, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + errorCode, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + finalSize, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + + var reliableSize uint64 + if isResetStreamAt { + reliableSize, l, err = quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + } + if reliableSize > finalSize { + return nil, 0, fmt.Errorf("RESET_STREAM_AT: reliable size can't be larger than final size (%d vs %d)", reliableSize, finalSize) + } + + return &ResetStreamFrame{ + StreamID: protocol.StreamID(streamID), + ErrorCode: qerr.StreamErrorCode(errorCode), + FinalSize: protocol.ByteCount(finalSize), + ReliableSize: protocol.ByteCount(reliableSize), + }, startLen - len(b), nil +} + +func (f *ResetStreamFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + if f.ReliableSize == 0 { + b = quicvarint.Append(b, uint64(FrameTypeResetStream)) + } else { + b = quicvarint.Append(b, uint64(FrameTypeResetStreamAt)) + } + b = quicvarint.Append(b, uint64(f.StreamID)) + b = quicvarint.Append(b, uint64(f.ErrorCode)) + b = quicvarint.Append(b, uint64(f.FinalSize)) + if f.ReliableSize > 0 { + b = quicvarint.Append(b, uint64(f.ReliableSize)) + } + return b, nil +} + +// Length of a written frame +func (f *ResetStreamFrame) Length(protocol.Version) protocol.ByteCount { + size := 1 // the frame type for both RESET_STREAM and RESET_STREAM_AT fits into 1 byte + if f.ReliableSize > 0 { + size += quicvarint.Len(uint64(f.ReliableSize)) + } + return protocol.ByteCount(size + quicvarint.Len(uint64(f.StreamID)) + quicvarint.Len(uint64(f.ErrorCode)) + quicvarint.Len(uint64(f.FinalSize))) +} diff --git a/third_party/quic-go/internal/wire/reset_stream_frame_test.go b/third_party/quic-go/internal/wire/reset_stream_frame_test.go new file mode 100644 index 0000000..17e9349 --- /dev/null +++ b/third_party/quic-go/internal/wire/reset_stream_frame_test.go @@ -0,0 +1,104 @@ +package wire + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + + "github.com/stretchr/testify/require" +) + +func TestParseResetStream(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // stream ID + data = append(data, encodeVarInt(0x1337)...) // error code + data = append(data, encodeVarInt(0x987654321)...) // byte offset + frame, l, err := parseResetStreamFrame(data, false, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0xdeadbeef), frame.StreamID) + require.Equal(t, protocol.ByteCount(0x987654321), frame.FinalSize) + require.Equal(t, qerr.StreamErrorCode(0x1337), frame.ErrorCode) + require.Equal(t, len(data), l) +} + +func TestParseResetStreamAt(t *testing.T) { + data := encodeVarInt(0xabcdef12) // stream ID + data = append(data, encodeVarInt(0x2468)...) // error code + data = append(data, encodeVarInt(0x123456789)...) // byte offset + data = append(data, encodeVarInt(0x789abc)...) // reliable size + frame, l, err := parseResetStreamFrame(data, true, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0xabcdef12), frame.StreamID) + require.Equal(t, protocol.ByteCount(0x123456789), frame.FinalSize) + require.Equal(t, qerr.StreamErrorCode(0x2468), frame.ErrorCode) + require.Equal(t, protocol.ByteCount(0x789abc), frame.ReliableSize) + require.Equal(t, len(data), l) +} + +func TestParseResetStreamAtSizeTooLarge(t *testing.T) { + data := encodeVarInt(0xabcdef12) // stream ID + data = append(data, encodeVarInt(0x2468)...) // error code + data = append(data, encodeVarInt(1000)...) // byte offset + data = append(data, encodeVarInt(1001)...) // reliable size + _, _, err := parseResetStreamFrame(data, true, protocol.Version1) + require.EqualError(t, err, "RESET_STREAM_AT: reliable size can't be larger than final size (1001 vs 1000)") +} + +func TestParseResetStreamErrorsOnEOFs(t *testing.T) { + t.Run("RESET_STREAM", func(t *testing.T) { + testParseResetStreamErrorsOnEOFs(t, false) + }) + t.Run("RESET_STREAM_AT", func(t *testing.T) { + testParseResetStreamErrorsOnEOFs(t, true) + }) +} + +func testParseResetStreamErrorsOnEOFs(t *testing.T, isResetStreamAt bool) { + data := encodeVarInt(0xdeadbeef) // stream ID + data = append(data, encodeVarInt(0x1337)...) // error code + data = append(data, encodeVarInt(0x987654321)...) // byte offset + if isResetStreamAt { + data = append(data, encodeVarInt(0x123456)...) // reliable size + } + _, l, err := parseResetStreamFrame(data, isResetStreamAt, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseResetStreamFrame(data[:i], isResetStreamAt, protocol.Version1) + require.Error(t, err) + } +} + +func TestWriteResetStream(t *testing.T) { + frame := ResetStreamFrame{ + StreamID: 0x1337, + FinalSize: 0x11223344decafbad, + ErrorCode: 0xcafe, + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeResetStream)} + expected = append(expected, encodeVarInt(0x1337)...) + expected = append(expected, encodeVarInt(0xcafe)...) + expected = append(expected, encodeVarInt(0x11223344decafbad)...) + require.Equal(t, expected, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} + +func TestWriteResetStreamAt(t *testing.T) { + frame := ResetStreamFrame{ + StreamID: 1337, + FinalSize: 42, + ErrorCode: 0xcafe, + ReliableSize: 12, + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeResetStreamAt)} + expected = append(expected, encodeVarInt(1337)...) + expected = append(expected, encodeVarInt(0xcafe)...) + expected = append(expected, encodeVarInt(42)...) + expected = append(expected, encodeVarInt(12)...) + require.Equal(t, expected, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/retire_connection_id_frame.go b/third_party/quic-go/internal/wire/retire_connection_id_frame.go new file mode 100644 index 0000000..635fff2 --- /dev/null +++ b/third_party/quic-go/internal/wire/retire_connection_id_frame.go @@ -0,0 +1,30 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A RetireConnectionIDFrame is a RETIRE_CONNECTION_ID frame +type RetireConnectionIDFrame struct { + SequenceNumber uint64 +} + +func parseRetireConnectionIDFrame(b []byte, _ protocol.Version) (*RetireConnectionIDFrame, int, error) { + seq, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + return &RetireConnectionIDFrame{SequenceNumber: seq}, l, nil +} + +func (f *RetireConnectionIDFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeRetireConnectionID)) + b = quicvarint.Append(b, f.SequenceNumber) + return b, nil +} + +// Length of a written frame +func (f *RetireConnectionIDFrame) Length(protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(f.SequenceNumber)) +} diff --git a/third_party/quic-go/internal/wire/retire_connection_id_frame_test.go b/third_party/quic-go/internal/wire/retire_connection_id_frame_test.go new file mode 100644 index 0000000..62603be --- /dev/null +++ b/third_party/quic-go/internal/wire/retire_connection_id_frame_test.go @@ -0,0 +1,39 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseRetireConnectionID(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // sequence number + frame, l, err := parseRetireConnectionIDFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, uint64(0xdeadbeef), frame.SequenceNumber) + require.Equal(t, len(data), l) +} + +func TestParseRetireConnectionIDErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // sequence number + _, l, err := parseRetireConnectionIDFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseRetireConnectionIDFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteRetireConnectionID(t *testing.T) { + frame := &RetireConnectionIDFrame{SequenceNumber: 0x1337} + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeRetireConnectionID)} + expected = append(expected, encodeVarInt(0x1337)...) + require.Equal(t, expected, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/short_header.go b/third_party/quic-go/internal/wire/short_header.go new file mode 100644 index 0000000..f0fe6e0 --- /dev/null +++ b/third_party/quic-go/internal/wire/short_header.go @@ -0,0 +1,62 @@ +package wire + +import ( + "errors" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +// ParseShortHeader parses a short header packet. +// It must be called after header protection was removed. +// Otherwise, the check for the reserved bits will (most likely) fail. +func ParseShortHeader(data []byte, connIDLen int) (length int, _ protocol.PacketNumber, _ protocol.PacketNumberLen, _ protocol.KeyPhaseBit, _ error) { + if len(data) == 0 { + return 0, 0, 0, 0, io.EOF + } + if data[0]&0x80 > 0 { + return 0, 0, 0, 0, errors.New("not a short header packet") + } + if data[0]&0x40 == 0 { + return 0, 0, 0, 0, errors.New("not a QUIC packet") + } + pnLen := protocol.PacketNumberLen(data[0]&0b11) + 1 + if len(data) < 1+int(pnLen)+connIDLen { + return 0, 0, 0, 0, io.EOF + } + + pos := 1 + connIDLen + pn, err := readPacketNumber(data[pos:], pnLen) + if err != nil { + return 0, 0, 0, 0, err + } + kp := protocol.KeyPhaseZero + if data[0]&0b100 > 0 { + kp = protocol.KeyPhaseOne + } + + if data[0]&0x18 != 0 { + err = ErrInvalidReservedBits + } + return 1 + connIDLen + int(pnLen), pn, pnLen, kp, err +} + +// AppendShortHeader writes a short header. +func AppendShortHeader(b []byte, connID protocol.ConnectionID, pn protocol.PacketNumber, pnLen protocol.PacketNumberLen, kp protocol.KeyPhaseBit) ([]byte, error) { + typeByte := 0x40 | uint8(pnLen-1) + if kp == protocol.KeyPhaseOne { + typeByte |= byte(1 << 2) + } + b = append(b, typeByte) + b = append(b, connID.Bytes()...) + return appendPacketNumber(b, pn, pnLen) +} + +func ShortHeaderLen(dest protocol.ConnectionID, pnLen protocol.PacketNumberLen) protocol.ByteCount { + return 1 + protocol.ByteCount(dest.Len()) + protocol.ByteCount(pnLen) +} + +func LogShortHeader(logger utils.Logger, dest protocol.ConnectionID, pn protocol.PacketNumber, pnLen protocol.PacketNumberLen, kp protocol.KeyPhaseBit) { + logger.Debugf("\tShort Header{DestConnectionID: %s, PacketNumber: %d, PacketNumberLen: %d, KeyPhase: %s}", dest, pn, pnLen, kp) +} diff --git a/third_party/quic-go/internal/wire/short_header_test.go b/third_party/quic-go/internal/wire/short_header_test.go new file mode 100644 index 0000000..b44d50a --- /dev/null +++ b/third_party/quic-go/internal/wire/short_header_test.go @@ -0,0 +1,106 @@ +package wire + +import ( + "bytes" + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseShortHeader(t *testing.T) { + data := []byte{ + 0b01000110, + 0xde, 0xad, 0xbe, 0xef, + 0x13, 0x37, 0x99, + } + l, pn, pnLen, kp, err := ParseShortHeader(data, 4) + require.NoError(t, err) + require.Equal(t, len(data), l) + require.Equal(t, protocol.KeyPhaseOne, kp) + require.Equal(t, protocol.PacketNumber(0x133799), pn) + require.Equal(t, protocol.PacketNumberLen3, pnLen) +} + +func TestParseShortHeaderNoQUICBit(t *testing.T) { + data := []byte{ + 0b00000101, + 0xde, 0xad, 0xbe, 0xef, + 0x13, 0x37, + } + _, _, _, _, err := ParseShortHeader(data, 4) + require.EqualError(t, err, "not a QUIC packet") +} + +func TestParseShortHeaderReservedBitsSet(t *testing.T) { + data := []byte{ + 0b01010101, + 0xde, 0xad, 0xbe, 0xef, + 0x13, 0x37, + } + _, pn, _, _, err := ParseShortHeader(data, 4) + require.EqualError(t, err, ErrInvalidReservedBits.Error()) + require.Equal(t, protocol.PacketNumber(0x1337), pn) +} + +func TestParseShortHeaderErrorsWhenPassedLongHeaderPacket(t *testing.T) { + _, _, _, _, err := ParseShortHeader([]byte{0x80}, 4) + require.EqualError(t, err, "not a short header packet") +} + +func TestParseShortHeaderErrorsOnEOF(t *testing.T) { + data := []byte{ + 0b01000110, + 0xde, 0xad, 0xbe, 0xef, + 0x13, 0x37, 0x99, + } + _, _, _, _, err := ParseShortHeader(data, 4) + require.NoError(t, err) + for i := range data { + _, _, _, _, err := ParseShortHeader(data[:i], 4) + require.EqualError(t, err, io.EOF.Error()) + } +} + +func TestShortHeaderLen(t *testing.T) { + require.Equal(t, protocol.ByteCount(8), ShortHeaderLen(protocol.ParseConnectionID([]byte{1, 2, 3, 4}), protocol.PacketNumberLen3)) + require.Equal(t, protocol.ByteCount(2), ShortHeaderLen(protocol.ParseConnectionID([]byte{}), protocol.PacketNumberLen1)) +} + +func TestWriteShortHeaderPacket(t *testing.T) { + connID := protocol.ParseConnectionID([]byte{1, 2, 3, 4}) + b, err := AppendShortHeader(nil, connID, 1337, 4, protocol.KeyPhaseOne) + require.NoError(t, err) + l, pn, pnLen, kp, err := ParseShortHeader(b, 4) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumber(1337), pn) + require.Equal(t, protocol.PacketNumberLen4, pnLen) + require.Equal(t, protocol.KeyPhaseOne, kp) + require.Equal(t, len(b), l) +} + +func TestLogShortHeaderWithConnectionID(t *testing.T) { + buf := &bytes.Buffer{} + logger := setupLogTest(t, buf) + + connID := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37}) + LogShortHeader(logger, connID, 1337, protocol.PacketNumberLen4, protocol.KeyPhaseOne) + require.Contains(t, buf.String(), "Short Header{DestConnectionID: deadbeefcafe1337, PacketNumber: 1337, PacketNumberLen: 4, KeyPhase: 1}") +} + +func BenchmarkWriteShortHeader(b *testing.B) { + b.ReportAllocs() + buf := make([]byte, 100) + connID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6}) + + for b.Loop() { + var err error + buf, err = AppendShortHeader(buf, connID, 1337, protocol.PacketNumberLen4, protocol.KeyPhaseOne) + if err != nil { + b.Fatalf("failed to write short header: %s", err) + } + buf = buf[:0] + } +} diff --git a/third_party/quic-go/internal/wire/stop_sending_frame.go b/third_party/quic-go/internal/wire/stop_sending_frame.go new file mode 100644 index 0000000..0a50e54 --- /dev/null +++ b/third_party/quic-go/internal/wire/stop_sending_frame.go @@ -0,0 +1,45 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/quicvarint" +) + +// A StopSendingFrame is a STOP_SENDING frame +type StopSendingFrame struct { + StreamID protocol.StreamID + ErrorCode qerr.StreamErrorCode +} + +// parseStopSendingFrame parses a STOP_SENDING frame +func parseStopSendingFrame(b []byte, _ protocol.Version) (*StopSendingFrame, int, error) { + startLen := len(b) + streamID, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + errorCode, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + + return &StopSendingFrame{ + StreamID: protocol.StreamID(streamID), + ErrorCode: qerr.StreamErrorCode(errorCode), + }, startLen - len(b), nil +} + +// Length of a written frame +func (f *StopSendingFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.StreamID))+quicvarint.Len(uint64(f.ErrorCode))) +} + +func (f *StopSendingFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, byte(FrameTypeStopSending)) + b = quicvarint.Append(b, uint64(f.StreamID)) + b = quicvarint.Append(b, uint64(f.ErrorCode)) + return b, nil +} diff --git a/third_party/quic-go/internal/wire/stop_sending_frame_test.go b/third_party/quic-go/internal/wire/stop_sending_frame_test.go new file mode 100644 index 0000000..5f2496d --- /dev/null +++ b/third_party/quic-go/internal/wire/stop_sending_frame_test.go @@ -0,0 +1,47 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + + "github.com/stretchr/testify/require" +) + +func TestParseStopSending(t *testing.T) { + data := encodeVarInt(0xdecafbad) // stream ID + data = append(data, encodeVarInt(0x1337)...) // error code + frame, l, err := parseStopSendingFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0xdecafbad), frame.StreamID) + require.Equal(t, qerr.StreamErrorCode(0x1337), frame.ErrorCode) + require.Equal(t, len(data), l) +} + +func TestParseStopSendingErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdecafbad) // stream ID + data = append(data, encodeVarInt(0x123456)...) // error code + _, l, err := parseStopSendingFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseStopSendingFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteStopSendingFrame(t *testing.T) { + frame := &StopSendingFrame{ + StreamID: 0xdeadbeefcafe, + ErrorCode: 0xdecafbad, + } + b, err := frame.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeStopSending)} + expected = append(expected, encodeVarInt(0xdeadbeefcafe)...) + expected = append(expected, encodeVarInt(0xdecafbad)...) + require.Equal(t, expected, b) + require.Len(t, b, int(frame.Length(protocol.Version1))) +} diff --git a/third_party/quic-go/internal/wire/stream_data_blocked_frame.go b/third_party/quic-go/internal/wire/stream_data_blocked_frame.go new file mode 100644 index 0000000..1877901 --- /dev/null +++ b/third_party/quic-go/internal/wire/stream_data_blocked_frame.go @@ -0,0 +1,42 @@ +package wire + +import ( + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A StreamDataBlockedFrame is a STREAM_DATA_BLOCKED frame +type StreamDataBlockedFrame struct { + StreamID protocol.StreamID + MaximumStreamData protocol.ByteCount +} + +func parseStreamDataBlockedFrame(b []byte, _ protocol.Version) (*StreamDataBlockedFrame, int, error) { + startLen := len(b) + sid, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + offset, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + + return &StreamDataBlockedFrame{ + StreamID: protocol.StreamID(sid), + MaximumStreamData: protocol.ByteCount(offset), + }, startLen - len(b) + l, nil +} + +func (f *StreamDataBlockedFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + b = append(b, 0x15) + b = quicvarint.Append(b, uint64(f.StreamID)) + b = quicvarint.Append(b, uint64(f.MaximumStreamData)) + return b, nil +} + +// Length of a written frame +func (f *StreamDataBlockedFrame) Length(protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.StreamID))+quicvarint.Len(uint64(f.MaximumStreamData))) +} diff --git a/third_party/quic-go/internal/wire/stream_data_blocked_frame_test.go b/third_party/quic-go/internal/wire/stream_data_blocked_frame_test.go new file mode 100644 index 0000000..a6e7e4a --- /dev/null +++ b/third_party/quic-go/internal/wire/stream_data_blocked_frame_test.go @@ -0,0 +1,46 @@ +package wire + +import ( + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseStreamDataBlocked(t *testing.T) { + data := encodeVarInt(0xdeadbeef) // stream ID + data = append(data, encodeVarInt(0xdecafbad)...) // offset + frame, l, err := parseStreamDataBlockedFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0xdeadbeef), frame.StreamID) + require.Equal(t, protocol.ByteCount(0xdecafbad), frame.MaximumStreamData) + require.Equal(t, len(data), l) +} + +func TestParseStreamDataBlockedErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0xdeadbeef) + data = append(data, encodeVarInt(0xc0010ff)...) + _, l, err := parseStreamDataBlockedFrame(data, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseStreamDataBlockedFrame(data[:i], protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestWriteStreamDataBlocked(t *testing.T) { + f := &StreamDataBlockedFrame{ + StreamID: 0xdecafbad, + MaximumStreamData: 0x1337, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeStreamDataBlocked)} + expected = append(expected, encodeVarInt(uint64(f.StreamID))...) + expected = append(expected, encodeVarInt(uint64(f.MaximumStreamData))...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} diff --git a/third_party/quic-go/internal/wire/stream_frame.go b/third_party/quic-go/internal/wire/stream_frame.go new file mode 100644 index 0000000..24b496f --- /dev/null +++ b/third_party/quic-go/internal/wire/stream_frame.go @@ -0,0 +1,191 @@ +package wire + +import ( + "errors" + "io" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A StreamFrame of QUIC +type StreamFrame struct { + StreamID protocol.StreamID + Offset protocol.ByteCount + Data []byte + Fin bool + DataLenPresent bool + + fromPool bool +} + +func ParseStreamFrame(b []byte, typ FrameType, _ protocol.Version) (*StreamFrame, int, error) { + startLen := len(b) + hasOffset := typ&0b100 > 0 + fin := typ&0b1 > 0 + hasDataLen := typ&0b10 > 0 + + streamID, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + var offset uint64 + if hasOffset { + offset, l, err = quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + } + + var dataLen uint64 + if hasDataLen { + var err error + var l int + dataLen, l, err = quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + b = b[l:] + if dataLen > uint64(len(b)) { + return nil, 0, io.EOF + } + } else { + // The rest of the packet is data + dataLen = uint64(len(b)) + } + + var frame *StreamFrame + if dataLen < protocol.MinStreamFrameBufferSize { + frame = &StreamFrame{} + if dataLen > 0 { + frame.Data = make([]byte, dataLen) + } + } else { + frame = GetStreamFrame() + // The STREAM frame can't be larger than the StreamFrame we obtained from the buffer, + // since those StreamFrames have a buffer length of the maximum packet size. + if dataLen > uint64(cap(frame.Data)) { + return nil, 0, io.EOF + } + frame.Data = frame.Data[:dataLen] + } + + frame.StreamID = protocol.StreamID(streamID) + frame.Offset = protocol.ByteCount(offset) + frame.Fin = fin + frame.DataLenPresent = hasDataLen + + if dataLen > 0 { + copy(frame.Data, b) + } + if frame.Offset+frame.DataLen() > protocol.MaxByteCount { + return nil, 0, errors.New("stream data overflows maximum offset") + } + return frame, startLen - len(b) + int(dataLen), nil +} + +func (f *StreamFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + if len(f.Data) == 0 && !f.Fin { + return nil, errors.New("StreamFrame: attempting to write empty frame without FIN") + } + + typ := byte(0x8) + if f.Fin { + typ ^= 0b1 + } + hasOffset := f.Offset != 0 + if f.DataLenPresent { + typ ^= 0b10 + } + if hasOffset { + typ ^= 0b100 + } + b = append(b, typ) + b = quicvarint.Append(b, uint64(f.StreamID)) + if hasOffset { + b = quicvarint.Append(b, uint64(f.Offset)) + } + if f.DataLenPresent { + b = quicvarint.Append(b, uint64(f.DataLen())) + } + b = append(b, f.Data...) + return b, nil +} + +// Length returns the total length of the STREAM frame +func (f *StreamFrame) Length(protocol.Version) protocol.ByteCount { + length := 1 + quicvarint.Len(uint64(f.StreamID)) + if f.Offset != 0 { + length += quicvarint.Len(uint64(f.Offset)) + } + if f.DataLenPresent { + length += quicvarint.Len(uint64(f.DataLen())) + } + return protocol.ByteCount(length) + f.DataLen() +} + +// DataLen gives the length of data in bytes +func (f *StreamFrame) DataLen() protocol.ByteCount { + return protocol.ByteCount(len(f.Data)) +} + +// MaxDataLen returns the maximum data length +// If 0 is returned, writing will fail (a STREAM frame must contain at least 1 byte of data). +func (f *StreamFrame) MaxDataLen(maxSize protocol.ByteCount, _ protocol.Version) protocol.ByteCount { + headerLen := 1 + protocol.ByteCount(quicvarint.Len(uint64(f.StreamID))) + if f.Offset != 0 { + headerLen += protocol.ByteCount(quicvarint.Len(uint64(f.Offset))) + } + if f.DataLenPresent { + // Pretend that the data size will be 1 byte. + // If it turns out that varint encoding the length will consume 2 bytes, we need to adjust the data length afterward + headerLen++ + } + if headerLen > maxSize { + return 0 + } + maxDataLen := maxSize - headerLen + if f.DataLenPresent && quicvarint.Len(uint64(maxDataLen)) != 1 { + maxDataLen-- + } + return maxDataLen +} + +// MaybeSplitOffFrame splits a frame such that it is not bigger than n bytes. +// It returns if the frame was actually split. +// The frame might not be split if: +// * the size is large enough to fit the whole frame +// * the size is too small to fit even a 1-byte frame. In that case, the frame returned is nil. +func (f *StreamFrame) MaybeSplitOffFrame(maxSize protocol.ByteCount, version protocol.Version) (*StreamFrame, bool /* was splitting required */) { + if maxSize >= f.Length(version) { + return nil, false + } + + n := f.MaxDataLen(maxSize, version) + if n == 0 { + return nil, true + } + + new := GetStreamFrame() + new.StreamID = f.StreamID + new.Offset = f.Offset + new.Fin = false + new.DataLenPresent = f.DataLenPresent + + // swap the data slices + new.Data, f.Data = f.Data, new.Data + new.fromPool, f.fromPool = f.fromPool, new.fromPool + + f.Data = f.Data[:protocol.ByteCount(len(new.Data))-n] + copy(f.Data, new.Data[n:]) + new.Data = new.Data[:n] + f.Offset += n + + return new, true +} + +func (f *StreamFrame) PutBack() { + putStreamFrame(f) +} diff --git a/third_party/quic-go/internal/wire/stream_frame_test.go b/third_party/quic-go/internal/wire/stream_frame_test.go new file mode 100644 index 0000000..7f3fee3 --- /dev/null +++ b/third_party/quic-go/internal/wire/stream_frame_test.go @@ -0,0 +1,379 @@ +package wire + +import ( + "bytes" + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseStreamFrameWithOffBit(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, encodeVarInt(0xdecafbad)...) // offset + data = append(data, []byte("foobar")...) + frame, l, err := ParseStreamFrame(data, 0x8^0x4, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0x12345), frame.StreamID) + require.Equal(t, []byte("foobar"), frame.Data) + require.False(t, frame.Fin) + require.Equal(t, protocol.ByteCount(0xdecafbad), frame.Offset) + require.Equal(t, len(data), l) +} + +func TestParseStreamFrameRespectsLEN(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, encodeVarInt(4)...) // data length + data = append(data, []byte("foobar")...) + frame, l, err := ParseStreamFrame(data, 0x8^0x2, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0x12345), frame.StreamID) + require.Equal(t, []byte("foob"), frame.Data) + require.False(t, frame.Fin) + require.Zero(t, frame.Offset) + require.Equal(t, len(data)-2, l) +} + +func TestParseStreamFrameWithFINBit(t *testing.T) { + data := encodeVarInt(9) // stream ID + data = append(data, []byte("foobar")...) + frame, l, err := ParseStreamFrame(data, 0x8^0x1, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(9), frame.StreamID) + require.Equal(t, []byte("foobar"), frame.Data) + require.True(t, frame.Fin) + require.Zero(t, frame.Offset) + require.Equal(t, len(data), l) +} + +func TestParseStreamFrameAllowsEmpty(t *testing.T) { + data := encodeVarInt(0x1337) // stream ID + data = append(data, encodeVarInt(0x12345)...) // offset + f, l, err := ParseStreamFrame(data, 0x8^0x4, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0x1337), f.StreamID) + require.Equal(t, protocol.ByteCount(0x12345), f.Offset) + require.Nil(t, f.Data) + require.False(t, f.Fin) + require.Equal(t, len(data), l) +} + +func TestParseStreamFrameRejectsOverflow(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, encodeVarInt(uint64(protocol.MaxByteCount-5))...) // offset + data = append(data, []byte("foobar")...) + _, _, err := ParseStreamFrame(data, 0x8^0x4, protocol.Version1) + require.EqualError(t, err, "stream data overflows maximum offset") +} + +func TestParseStreamFrameRejectsLongFrames(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, encodeVarInt(uint64(protocol.MaxPacketBufferSize)+1)...) // data length + data = append(data, make([]byte, protocol.MaxPacketBufferSize+1)...) + _, _, err := ParseStreamFrame(data, 0x8^0x2, protocol.Version1) + require.Equal(t, io.EOF, err) +} + +func TestParseStreamFrameRejectsFramesExceedingRemainingSize(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, encodeVarInt(7)...) // data length + data = append(data, []byte("foobar")...) + _, _, err := ParseStreamFrame(data, 0x8^0x2, protocol.Version1) + require.Equal(t, io.EOF, err) +} + +func TestParseStreamFrameErrorsOnEOFs(t *testing.T) { + typ := uint64(0x8 ^ 0x4 ^ 0x2) + data := encodeVarInt(0x12345) // stream ID + data = append(data, encodeVarInt(0xdecafbad)...) // offset + data = append(data, encodeVarInt(6)...) // data length + data = append(data, []byte("foobar")...) + _, _, err := ParseStreamFrame(data, FrameType(typ), protocol.Version1) + require.NoError(t, err) + for i := range data { + _, _, err = ParseStreamFrame(data[:i], FrameType(typ), protocol.Version1) + require.Error(t, err) + } +} + +func TestParseStreamUsesBufferForLongFrames(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, bytes.Repeat([]byte{'f'}, protocol.MinStreamFrameBufferSize)...) + frame, l, err := ParseStreamFrame(data, 0x8, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0x12345), frame.StreamID) + require.Equal(t, bytes.Repeat([]byte{'f'}, protocol.MinStreamFrameBufferSize), frame.Data) + require.Equal(t, protocol.ByteCount(protocol.MinStreamFrameBufferSize), frame.DataLen()) + require.False(t, frame.Fin) + require.True(t, frame.fromPool) + require.Equal(t, len(data), l) + require.NotPanics(t, frame.PutBack) +} + +func TestParseStreamDoesNotUseBufferForShortFrames(t *testing.T) { + data := encodeVarInt(0x12345) // stream ID + data = append(data, bytes.Repeat([]byte{'f'}, protocol.MinStreamFrameBufferSize-1)...) + frame, l, err := ParseStreamFrame(data, 0x8, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamID(0x12345), frame.StreamID) + require.Equal(t, bytes.Repeat([]byte{'f'}, protocol.MinStreamFrameBufferSize-1), frame.Data) + require.Equal(t, protocol.ByteCount(protocol.MinStreamFrameBufferSize-1), frame.DataLen()) + require.False(t, frame.Fin) + require.False(t, frame.fromPool) + require.Equal(t, len(data), l) + require.NotPanics(t, frame.PutBack) +} + +func TestWriteStreamFrameWithoutOffset(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + Data: []byte("foobar"), + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x8} + expected = append(expected, encodeVarInt(0x1337)...) // stream ID + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestWriteStreamFrameWithOffset(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + Offset: 0x123456, + Data: []byte("foobar"), + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x8 ^ 0x4} + expected = append(expected, encodeVarInt(0x1337)...) // stream ID + expected = append(expected, encodeVarInt(0x123456)...) // offset + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestWriteStreamFrameWithFIN(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + Offset: 0x123456, + Fin: true, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x8 ^ 0x4 ^ 0x1} + expected = append(expected, encodeVarInt(0x1337)...) // stream ID + expected = append(expected, encodeVarInt(0x123456)...) // offset + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestWriteStreamFrameWithDataLength(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + Data: []byte("foobar"), + DataLenPresent: true, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x8 ^ 0x2} + expected = append(expected, encodeVarInt(0x1337)...) // stream ID + expected = append(expected, encodeVarInt(6)...) // data length + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestWriteStreamFrameWithDataLengthAndOffset(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + Data: []byte("foobar"), + DataLenPresent: true, + Offset: 0x123456, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{0x8 ^ 0x4 ^ 0x2} + expected = append(expected, encodeVarInt(0x1337)...) // stream ID + expected = append(expected, encodeVarInt(0x123456)...) // offset + expected = append(expected, encodeVarInt(6)...) // data length + expected = append(expected, []byte("foobar")...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestWriteStreamFrameEmptyFrameWithoutFIN(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x42, + Offset: 0x1337, + } + _, err := f.Append(nil, protocol.Version1) + require.EqualError(t, err, "StreamFrame: attempting to write empty frame without FIN") +} + +func TestStreamMaxDataLength(t *testing.T) { + const maxSize = 3000 + data := make([]byte, maxSize) + f := &StreamFrame{ + StreamID: 0x1337, + Offset: 0xdeadbeef, + } + for i := 1; i < 3000; i++ { + f.Data = nil + maxDataLen := f.MaxDataLen(protocol.ByteCount(i), protocol.Version1) + if maxDataLen == 0 { // 0 means that no valid STREAM frame can be written + // check that writing a minimal size STREAM frame (i.e. with 1 byte data) is actually larger than the desired size + f.Data = []byte{0} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Greater(t, len(b), i) + continue + } + f.Data = data[:int(maxDataLen)] + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Equal(t, i, len(b)) + } +} + +func TestStreamMaxDataLengthWithDataLenPresent(t *testing.T) { + const maxSize = 3000 + data := make([]byte, maxSize) + f := &StreamFrame{ + StreamID: 0x1337, + Offset: 0xdeadbeef, + DataLenPresent: true, + } + var frameOneByteTooSmallCounter int + for i := 1; i < 3000; i++ { + f.Data = nil + maxDataLen := f.MaxDataLen(protocol.ByteCount(i), protocol.Version1) + if maxDataLen == 0 { // 0 means that no valid STREAM frame can be written + // check that writing a minimal size STREAM frame (i.e. with 1 byte data) is actually larger than the desired size + f.Data = []byte{0} + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Greater(t, len(b), i) + continue + } + f.Data = data[:int(maxDataLen)] + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + // There's *one* pathological case, where a data length of x can be encoded into 1 byte + // but a data lengths of x+1 needs 2 bytes + // In that case, it's impossible to create a STREAM frame of the desired size + if len(b) == i-1 { + frameOneByteTooSmallCounter++ + continue + } + require.Equal(t, i, len(b)) + } + require.Equal(t, 1, frameOneByteTooSmallCounter) +} + +func TestStreamSplitting(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + DataLenPresent: true, + Offset: 0x100, + Data: []byte("foobar"), + } + frame, needsSplit := f.MaybeSplitOffFrame(f.Length(protocol.Version1)-3, protocol.Version1) + require.True(t, needsSplit) + require.NotNil(t, frame) + require.True(t, f.DataLenPresent) + require.True(t, frame.DataLenPresent) + require.Equal(t, protocol.ByteCount(0x100), frame.Offset) + require.Equal(t, []byte("foo"), frame.Data) + require.Equal(t, protocol.ByteCount(0x100+3), f.Offset) + require.Equal(t, []byte("bar"), f.Data) +} + +func TestStreamSplittingNoSplitForShortFrame(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + DataLenPresent: true, + Offset: 0xdeadbeef, + Data: make([]byte, 100), + } + frame, needsSplit := f.MaybeSplitOffFrame(f.Length(protocol.Version1), protocol.Version1) + require.False(t, needsSplit) + require.Nil(t, frame) + require.Equal(t, protocol.ByteCount(100), f.DataLen()) + frame, needsSplit = f.MaybeSplitOffFrame(f.Length(protocol.Version1)-1, protocol.Version1) + require.True(t, needsSplit) + require.Equal(t, protocol.ByteCount(99), frame.DataLen()) + f.PutBack() +} + +func TestStreamSplittingPreservesFINBit(t *testing.T) { + f := &StreamFrame{ + StreamID: 0x1337, + Fin: true, + Offset: 0xdeadbeef, + Data: make([]byte, 100), + } + frame, needsSplit := f.MaybeSplitOffFrame(50, protocol.Version1) + require.True(t, needsSplit) + require.NotNil(t, frame) + require.Less(t, frame.Offset, f.Offset) + require.True(t, f.Fin) + require.False(t, frame.Fin) +} + +func TestStreamSplittingProducesCorrectLengthFramesWithoutDataLen(t *testing.T) { + const size = 1000 + f := &StreamFrame{ + StreamID: 0xdecafbad, + Offset: 0x1234, + Data: []byte{0}, + } + minFrameSize := f.Length(protocol.Version1) + for i := range minFrameSize { + frame, needsSplit := f.MaybeSplitOffFrame(i, protocol.Version1) + require.True(t, needsSplit) + require.Nil(t, frame) + } + for i := minFrameSize; i < size; i++ { + f.fromPool = false + f.Data = make([]byte, size) + frame, needsSplit := f.MaybeSplitOffFrame(i, protocol.Version1) + require.True(t, needsSplit) + require.Equal(t, i, frame.Length(protocol.Version1)) + } +} + +func TestStreamSplittingProducesCorrectLengthFramesWithDataLen(t *testing.T) { + const size = 1000 + f := &StreamFrame{ + StreamID: 0xdecafbad, + Offset: 0x1234, + DataLenPresent: true, + Data: []byte{0}, + } + minFrameSize := f.Length(protocol.Version1) + for i := range minFrameSize { + frame, needsSplit := f.MaybeSplitOffFrame(i, protocol.Version1) + require.True(t, needsSplit) + require.Nil(t, frame) + } + var frameOneByteTooSmallCounter int + for i := minFrameSize; i < size; i++ { + f.fromPool = false + f.Data = make([]byte, size) + newFrame, needsSplit := f.MaybeSplitOffFrame(i, protocol.Version1) + require.True(t, needsSplit) + // There's *one* pathological case, where a data length of x can be encoded into 1 byte + // but a data lengths of x+1 needs 2 bytes + // In that case, it's impossible to create a STREAM frame of the desired size + if newFrame.Length(protocol.Version1) == i-1 { + frameOneByteTooSmallCounter++ + continue + } + require.Equal(t, i, newFrame.Length(protocol.Version1)) + } + require.Equal(t, 1, frameOneByteTooSmallCounter) +} diff --git a/third_party/quic-go/internal/wire/streams_blocked_frame.go b/third_party/quic-go/internal/wire/streams_blocked_frame.go new file mode 100644 index 0000000..9a27f34 --- /dev/null +++ b/third_party/quic-go/internal/wire/streams_blocked_frame.go @@ -0,0 +1,50 @@ +package wire + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// A StreamsBlockedFrame is a STREAMS_BLOCKED frame +type StreamsBlockedFrame struct { + Type protocol.StreamType + StreamLimit protocol.StreamNum +} + +func parseStreamsBlockedFrame(b []byte, typ FrameType, _ protocol.Version) (*StreamsBlockedFrame, int, error) { + f := &StreamsBlockedFrame{} + //nolint:exhaustive // This will only be called with a BidiStreamBlockedFrameType or a UniStreamBlockedFrameType. + switch typ { + case FrameTypeBidiStreamBlocked: + f.Type = protocol.StreamTypeBidi + case FrameTypeUniStreamBlocked: + f.Type = protocol.StreamTypeUni + } + streamLimit, l, err := quicvarint.Parse(b) + if err != nil { + return nil, 0, replaceUnexpectedEOF(err) + } + f.StreamLimit = protocol.StreamNum(streamLimit) + if f.StreamLimit > protocol.MaxStreamCount { + return nil, 0, fmt.Errorf("%d exceeds the maximum stream count", f.StreamLimit) + } + return f, l, nil +} + +func (f *StreamsBlockedFrame) Append(b []byte, _ protocol.Version) ([]byte, error) { + switch f.Type { + case protocol.StreamTypeBidi: + b = append(b, byte(FrameTypeBidiStreamBlocked)) + case protocol.StreamTypeUni: + b = append(b, byte(FrameTypeUniStreamBlocked)) + } + b = quicvarint.Append(b, uint64(f.StreamLimit)) + return b, nil +} + +// Length of a written frame +func (f *StreamsBlockedFrame) Length(_ protocol.Version) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(f.StreamLimit))) +} diff --git a/third_party/quic-go/internal/wire/streams_blocked_frame_test.go b/third_party/quic-go/internal/wire/streams_blocked_frame_test.go new file mode 100644 index 0000000..79a2948 --- /dev/null +++ b/third_party/quic-go/internal/wire/streams_blocked_frame_test.go @@ -0,0 +1,117 @@ +package wire + +import ( + "fmt" + "io" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + + "github.com/stretchr/testify/require" +) + +func TestParseStreamsBlockedFrameBidirectional(t *testing.T) { + data := encodeVarInt(0x1337) + f, l, err := parseStreamsBlockedFrame(data, FrameTypeBidiStreamBlocked, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamTypeBidi, f.Type) + require.EqualValues(t, 0x1337, f.StreamLimit) + require.Equal(t, len(data), l) +} + +func TestParseStreamsBlockedFrameUnidirectional(t *testing.T) { + data := encodeVarInt(0x7331) + f, l, err := parseStreamsBlockedFrame(data, FrameTypeUniStreamBlocked, protocol.Version1) + require.NoError(t, err) + require.Equal(t, protocol.StreamTypeUni, f.Type) + require.EqualValues(t, 0x7331, f.StreamLimit) + require.Equal(t, len(data), l) +} + +func TestParseStreamsBlockedFrameErrorsOnEOFs(t *testing.T) { + data := encodeVarInt(0x12345678) + _, l, err := parseStreamsBlockedFrame(data, FrameTypeBidiStreamBlocked, protocol.Version1) + require.NoError(t, err) + require.Equal(t, len(data), l) + for i := range data { + _, _, err := parseStreamsBlockedFrame(data[:i], FrameTypeBidiStreamBlocked, protocol.Version1) + require.Equal(t, io.EOF, err) + } +} + +func TestParseStreamsBlockedFrameMaxStreamCount(t *testing.T) { + for _, streamType := range []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi} { + var streamTypeStr string + if streamType == protocol.StreamTypeUni { + streamTypeStr = "unidirectional" + } else { + streamTypeStr = "bidirectional" + } + t.Run(streamTypeStr, func(t *testing.T) { + f := &StreamsBlockedFrame{ + Type: streamType, + StreamLimit: protocol.MaxStreamCount, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + typ, l, err := quicvarint.Parse(b) + require.NoError(t, err) + b = b[l:] + frame, l, err := parseStreamsBlockedFrame(b, FrameType(typ), protocol.Version1) + require.NoError(t, err) + require.Equal(t, f, frame) + require.Equal(t, len(b), l) + }) + } +} + +func TestParseStreamsBlockedFrameErrorOnTooLargeStreamCount(t *testing.T) { + for _, streamType := range []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi} { + var streamTypeStr string + if streamType == protocol.StreamTypeUni { + streamTypeStr = "unidirectional" + } else { + streamTypeStr = "bidirectional" + } + t.Run(streamTypeStr, func(t *testing.T) { + f := &StreamsBlockedFrame{ + Type: streamType, + StreamLimit: protocol.MaxStreamCount + 1, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + typ, l, err := quicvarint.Parse(b) + require.NoError(t, err) + b = b[l:] + _, _, err = parseStreamsBlockedFrame(b, FrameType(typ), protocol.Version1) + require.EqualError(t, err, fmt.Sprintf("%d exceeds the maximum stream count", protocol.MaxStreamCount+1)) + }) + } +} + +func TestWriteStreamsBlockedFrameBidirectional(t *testing.T) { + f := StreamsBlockedFrame{ + Type: protocol.StreamTypeBidi, + StreamLimit: 0xdeadbeefcafe, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeBidiStreamBlocked)} + expected = append(expected, encodeVarInt(0xdeadbeefcafe)...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} + +func TestWriteStreamsBlockedFrameUnidirectional(t *testing.T) { + f := StreamsBlockedFrame{ + Type: protocol.StreamTypeUni, + StreamLimit: 0xdeadbeefcafe, + } + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + expected := []byte{byte(FrameTypeUniStreamBlocked)} + expected = append(expected, encodeVarInt(0xdeadbeefcafe)...) + require.Equal(t, expected, b) + require.Equal(t, int(f.Length(protocol.Version1)), len(b)) +} diff --git a/third_party/quic-go/internal/wire/test_helpers_test.go b/third_party/quic-go/internal/wire/test_helpers_test.go new file mode 100644 index 0000000..abfd166 --- /dev/null +++ b/third_party/quic-go/internal/wire/test_helpers_test.go @@ -0,0 +1,32 @@ +package wire + +import ( + "bytes" + "encoding/binary" + "log" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/quicvarint" +) + +func encodeVarInt(i uint64) []byte { + return quicvarint.Append(nil, i) +} + +func appendVersion(data []byte, v protocol.Version) []byte { + offset := len(data) + data = append(data, []byte{0, 0, 0, 0}...) + binary.BigEndian.PutUint32(data[offset:], uint32(v)) + return data +} + +func setupLogTest(t *testing.T, buf *bytes.Buffer) utils.Logger { + logger := utils.DefaultLogger + logger.SetLogLevel(utils.LogLevelDebug) + originalOutput := log.Writer() + log.SetOutput(buf) + t.Cleanup(func() { log.SetOutput(originalOutput) }) + return logger +} diff --git a/third_party/quic-go/internal/wire/transport_parameter_test.go b/third_party/quic-go/internal/wire/transport_parameter_test.go new file mode 100644 index 0000000..c83dfef --- /dev/null +++ b/third_party/quic-go/internal/wire/transport_parameter_test.go @@ -0,0 +1,1172 @@ +package wire + +import ( + "bytes" + "crypto/rand" + "fmt" + "math" + mrand "math/rand/v2" + "net/netip" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/quicvarint" + + ossfuzzseeds "github.com/quic-go/go-ossfuzz-seeds" + + "github.com/stretchr/testify/require" +) + +func getRandomValueUpTo(max uint64) uint64 { + maxVals := []uint64{math.MaxUint8 / 4, math.MaxUint16 / 4, math.MaxUint32 / 4, math.MaxUint64 / 4} + return mrand.Uint64N(min(max, maxVals[mrand.IntN(4)])) +} + +func getRandomValue() uint64 { return getRandomValueUpTo(quicvarint.Max) } + +func appendInitialSourceConnectionID(b []byte) []byte { + b = quicvarint.Append(b, uint64(initialSourceConnectionIDParameterID)) + b = quicvarint.Append(b, 6) + return append(b, []byte("foobar")...) +} + +func TestTransportParametersStringRepresentation(t *testing.T) { + rcid := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}) + minAckDelay := 42 * time.Millisecond + p := &TransportParameters{ + InitialMaxStreamDataBidiLocal: 1234, + InitialMaxStreamDataBidiRemote: 2345, + InitialMaxStreamDataUni: 3456, + InitialMaxData: 4567, + MaxBidiStreamNum: 1337, + MaxUniStreamNum: 7331, + MaxIdleTimeout: 42 * time.Second, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + RetrySourceConnectionID: &rcid, + AckDelayExponent: 14, + MaxAckDelay: 37 * time.Millisecond, + StatelessResetToken: &protocol.StatelessResetToken{0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff, 0x00}, + ActiveConnectionIDLimit: 123, + MaxDatagramFrameSize: 876, + EnableResetStreamAt: true, + MinAckDelay: &minAckDelay, + } + expected := "&wire.TransportParameters{OriginalDestinationConnectionID: deadbeef, InitialSourceConnectionID: decafbad, RetrySourceConnectionID: deadc0de, InitialMaxStreamDataBidiLocal: 1234, InitialMaxStreamDataBidiRemote: 2345, InitialMaxStreamDataUni: 3456, InitialMaxData: 4567, MaxBidiStreamNum: 1337, MaxUniStreamNum: 7331, MaxIdleTimeout: 42s, AckDelayExponent: 14, MaxAckDelay: 37ms, ActiveConnectionIDLimit: 123, StatelessResetToken: 0x112233445566778899aabbccddeeff00, MaxDatagramFrameSize: 876, EnableResetStreamAt: true, MinAckDelay: 42ms}" + require.Equal(t, expected, p.String()) +} + +func TestTransportParametersStringRepresentationWithoutOptionalFields(t *testing.T) { + p := &TransportParameters{ + InitialMaxStreamDataBidiLocal: 1234, + InitialMaxStreamDataBidiRemote: 2345, + InitialMaxStreamDataUni: 3456, + InitialMaxData: 4567, + MaxBidiStreamNum: 1337, + MaxUniStreamNum: 7331, + MaxIdleTimeout: 42 * time.Second, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{}), + AckDelayExponent: 14, + MaxAckDelay: 37 * time.Second, + ActiveConnectionIDLimit: 89, + MaxDatagramFrameSize: protocol.InvalidByteCount, + } + expected := "&wire.TransportParameters{OriginalDestinationConnectionID: deadbeef, InitialSourceConnectionID: (empty), InitialMaxStreamDataBidiLocal: 1234, InitialMaxStreamDataBidiRemote: 2345, InitialMaxStreamDataUni: 3456, InitialMaxData: 4567, MaxBidiStreamNum: 1337, MaxUniStreamNum: 7331, MaxIdleTimeout: 42s, AckDelayExponent: 14, MaxAckDelay: 37s, ActiveConnectionIDLimit: 89, EnableResetStreamAt: false}" + require.Equal(t, expected, p.String()) +} + +func TestMarshalAndUnmarshalTransportParameters(t *testing.T) { + var token protocol.StatelessResetToken + rand.Read(token[:]) + rcid := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}) + minAckDelay := 42 * time.Millisecond + params := &TransportParameters{ + InitialMaxStreamDataBidiLocal: protocol.ByteCount(getRandomValue()), + InitialMaxStreamDataBidiRemote: protocol.ByteCount(getRandomValue()), + InitialMaxStreamDataUni: protocol.ByteCount(getRandomValue()), + InitialMaxData: protocol.ByteCount(getRandomValue()), + MaxIdleTimeout: 0xcafe * time.Second, + MaxBidiStreamNum: protocol.StreamNum(getRandomValueUpTo(uint64(protocol.MaxStreamCount))), + MaxUniStreamNum: protocol.StreamNum(getRandomValueUpTo(uint64(protocol.MaxStreamCount))), + DisableActiveMigration: true, + StatelessResetToken: &token, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + RetrySourceConnectionID: &rcid, + AckDelayExponent: 13, + MaxAckDelay: 42 * time.Millisecond, + ActiveConnectionIDLimit: 2 + getRandomValueUpTo(quicvarint.Max-2), + MaxUDPPayloadSize: 1200 + protocol.ByteCount(getRandomValueUpTo(quicvarint.Max-1200)), + MaxDatagramFrameSize: protocol.ByteCount(getRandomValue()), + EnableResetStreamAt: getRandomValue()%2 == 0, + MinAckDelay: &minAckDelay, + } + data := params.Marshal(protocol.PerspectiveServer) + + p := &TransportParameters{} + require.NoError(t, p.Unmarshal(data, protocol.PerspectiveServer)) + require.Equal(t, params.InitialMaxStreamDataBidiLocal, p.InitialMaxStreamDataBidiLocal) + require.Equal(t, params.InitialMaxStreamDataBidiRemote, p.InitialMaxStreamDataBidiRemote) + require.Equal(t, params.InitialMaxStreamDataUni, p.InitialMaxStreamDataUni) + require.Equal(t, params.InitialMaxData, p.InitialMaxData) + require.Equal(t, params.MaxUniStreamNum, p.MaxUniStreamNum) + require.Equal(t, params.MaxBidiStreamNum, p.MaxBidiStreamNum) + require.Equal(t, params.MaxIdleTimeout, p.MaxIdleTimeout) + require.Equal(t, params.DisableActiveMigration, p.DisableActiveMigration) + require.Equal(t, params.StatelessResetToken, p.StatelessResetToken) + require.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), p.OriginalDestinationConnectionID) + require.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), p.InitialSourceConnectionID) + require.Equal(t, &rcid, p.RetrySourceConnectionID) + require.Equal(t, uint8(13), p.AckDelayExponent) + require.Equal(t, 42*time.Millisecond, p.MaxAckDelay) + require.Equal(t, params.ActiveConnectionIDLimit, p.ActiveConnectionIDLimit) + require.Equal(t, params.MaxUDPPayloadSize, p.MaxUDPPayloadSize) + require.Equal(t, params.MaxDatagramFrameSize, p.MaxDatagramFrameSize) + require.Equal(t, params.EnableResetStreamAt, p.EnableResetStreamAt) + require.NotNil(t, p.MinAckDelay) + require.Equal(t, minAckDelay, *p.MinAckDelay) +} + +func TestResetStreamAtTransportParameterCodepoints(t *testing.T) { + for _, tc := range []struct { + name string + ids []transportParameterID + }{ + {name: "current", ids: []transportParameterID{resetStreamAtParameterID}}, + {name: "legacy", ids: []transportParameterID{legacyResetStreamAtParameterID}}, + {name: "both", ids: []transportParameterID{resetStreamAtParameterID, legacyResetStreamAtParameterID}}, + } { + t.Run(tc.name, func(t *testing.T) { + var data []byte + for _, id := range tc.ids { + data = quicvarint.Append(data, uint64(id)) + data = quicvarint.Append(data, 0) + } + data = appendInitialSourceConnectionID(data) + + var p TransportParameters + require.NoError(t, p.Unmarshal(data, protocol.PerspectiveClient)) + require.True(t, p.EnableResetStreamAt) + }) + } +} + +func TestMarshalAdditionalTransportParameters(t *testing.T) { + origAdditionalTransportParametersClient := AdditionalTransportParametersClient + t.Cleanup(func() { + AdditionalTransportParametersClient = origAdditionalTransportParametersClient + }) + AdditionalTransportParametersClient = map[uint64][]byte{1337: []byte("foobar")} + + result := quicvarint.Append([]byte{}, 1337) + result = quicvarint.Append(result, 6) + result = append(result, []byte("foobar")...) + + params := &TransportParameters{} + require.True(t, bytes.Contains(params.Marshal(protocol.PerspectiveClient), result)) + require.False(t, bytes.Contains(params.Marshal(protocol.PerspectiveServer), result)) +} + +func TestMarshalRetrySourceConnectionID(t *testing.T) { + // no retry source connection ID + data := (&TransportParameters{ + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: 2, + }).Marshal(protocol.PerspectiveServer) + var p TransportParameters + require.NoError(t, p.Unmarshal(data, protocol.PerspectiveServer)) + require.Nil(t, p.RetrySourceConnectionID) + + // zero-length retry source connection ID + rcid := protocol.ParseConnectionID([]byte{}) + data = (&TransportParameters{ + RetrySourceConnectionID: &rcid, + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: 2, + }).Marshal(protocol.PerspectiveServer) + p = TransportParameters{} + require.NoError(t, p.Unmarshal(data, protocol.PerspectiveServer)) + require.NotNil(t, p.RetrySourceConnectionID) + require.Zero(t, p.RetrySourceConnectionID.Len()) +} + +func TestTransportParameterNoMaxAckDelayIfDefault(t *testing.T) { + const num = 1000 + var defaultLen, dataLen int + maxAckDelay := protocol.DefaultMaxAckDelay + time.Millisecond + for range num { + dataDefault := (&TransportParameters{ + MaxAckDelay: protocol.DefaultMaxAckDelay, + StatelessResetToken: &protocol.StatelessResetToken{}, + }).Marshal(protocol.PerspectiveServer) + defaultLen += len(dataDefault) + data := (&TransportParameters{ + MaxAckDelay: maxAckDelay, + StatelessResetToken: &protocol.StatelessResetToken{}, + }).Marshal(protocol.PerspectiveServer) + dataLen += len(data) + } + entryLen := quicvarint.Len(uint64(ackDelayExponentParameterID)) + + quicvarint.Len(uint64(quicvarint.Len(uint64(maxAckDelay.Milliseconds())))) + + quicvarint.Len(uint64(maxAckDelay.Milliseconds())) + require.InDelta(t, float32(defaultLen)/num+float32(entryLen), float32(dataLen)/num, 1) +} + +func TestTransportParameterNoAckDelayExponentIfDefault(t *testing.T) { + const num = 1000 + var defaultLen, dataLen int + for range num { + dataDefault := (&TransportParameters{ + AckDelayExponent: protocol.DefaultAckDelayExponent, + StatelessResetToken: &protocol.StatelessResetToken{}, + }).Marshal(protocol.PerspectiveServer) + defaultLen += len(dataDefault) + data := (&TransportParameters{ + AckDelayExponent: protocol.DefaultAckDelayExponent + 1, + StatelessResetToken: &protocol.StatelessResetToken{}, + }).Marshal(protocol.PerspectiveServer) + dataLen += len(data) + } + entryLen := quicvarint.Len(uint64(ackDelayExponentParameterID)) + + quicvarint.Len(uint64(quicvarint.Len(protocol.DefaultAckDelayExponent+1))) + + quicvarint.Len(protocol.DefaultAckDelayExponent+1) + require.InDelta(t, float32(defaultLen)/num+float32(entryLen), float32(dataLen)/num, 1) +} + +func TestTransportParameterSetsDefaultValuesWhenNotSent(t *testing.T) { + data := (&TransportParameters{ + AckDelayExponent: protocol.DefaultAckDelayExponent, + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: protocol.DefaultActiveConnectionIDLimit, + }).Marshal(protocol.PerspectiveServer) + p := &TransportParameters{} + require.NoError(t, p.Unmarshal(data, protocol.PerspectiveServer)) + require.EqualValues(t, protocol.DefaultAckDelayExponent, p.AckDelayExponent) + require.EqualValues(t, protocol.DefaultActiveConnectionIDLimit, p.ActiveConnectionIDLimit) +} + +func TestTransportParameterErrors(t *testing.T) { + tests := []struct { + name string + params *TransportParameters + perspective protocol.Perspective + data []byte + expectedErrMsg string + }{ + { + name: "invalid stateless reset token length", + data: func() []byte { + b := quicvarint.Append(nil, uint64(statelessResetTokenParameterID)) + b = quicvarint.Append(b, 15) + return append(b, make([]byte, 15)...) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "wrong length for stateless_reset_token: 15 (expected 16)", + }, + { + name: "small max UDP payload size", + data: func() []byte { + b := quicvarint.Append(nil, uint64(maxUDPPayloadSizeParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(1199))) + return quicvarint.Append(b, 1199) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "invalid value for max_udp_payload_size: 1199 (minimum 1200)", + }, + { + name: "active connection ID limit too small", + params: &TransportParameters{ + ActiveConnectionIDLimit: 1, + StatelessResetToken: &protocol.StatelessResetToken{}, + }, + perspective: protocol.PerspectiveServer, + expectedErrMsg: "invalid value for active_connection_id_limit: 1 (minimum 2)", + }, + { + name: "ack delay exponent too large", + params: &TransportParameters{ + AckDelayExponent: 21, + StatelessResetToken: &protocol.StatelessResetToken{}, + }, + perspective: protocol.PerspectiveServer, + expectedErrMsg: "invalid value for ack_delay_exponent: 21 (maximum 20)", + }, + { + name: "disable active migration has content", + data: func() []byte { + b := quicvarint.Append(nil, uint64(disableActiveMigrationParameterID)) + b = quicvarint.Append(b, 6) + return append(b, []byte("foobar")...) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "wrong length for disable_active_migration: 6 (expected empty)", + }, + { + name: "server doesn't set original destination connection ID", + data: func() []byte { + b := quicvarint.Append(nil, uint64(statelessResetTokenParameterID)) + b = quicvarint.Append(b, 16) + b = append(b, make([]byte, 16)...) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "missing original_destination_connection_id", + }, + { + name: "initial source connection ID is missing", + data: []byte{}, + perspective: protocol.PerspectiveClient, + expectedErrMsg: "missing initial_source_connection_id", + }, + { + name: "max ack delay is too large", + params: &TransportParameters{ + MaxAckDelay: 1 << 14 * time.Millisecond, + StatelessResetToken: &protocol.StatelessResetToken{}, + }, + perspective: protocol.PerspectiveServer, + expectedErrMsg: "invalid value for max_ack_delay: 16384ms (maximum 16383ms)", + }, + { + name: "varint value has wrong length", + data: func() []byte { + b := quicvarint.Append(nil, uint64(initialMaxStreamDataBidiLocalParameterID)) + b = quicvarint.Append(b, 2) + val := uint64(0xdeadbeef) + b = quicvarint.Append(b, val) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: fmt.Sprintf("inconsistent transport parameter length for transport parameter %#x", initialMaxStreamDataBidiLocalParameterID), + }, + { + name: "initial max streams bidi is too large", + data: func() []byte { + b := quicvarint.Append(nil, uint64(initialMaxStreamsBidiParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(uint64(protocol.MaxStreamCount+1)))) + b = quicvarint.Append(b, uint64(protocol.MaxStreamCount+1)) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "initial_max_streams_bidi too large: 1152921504606846977 (maximum 1152921504606846976)", + }, + { + name: "initial max streams uni is too large", + data: func() []byte { + b := quicvarint.Append(nil, uint64(initialMaxStreamsUniParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(uint64(protocol.MaxStreamCount+1)))) + b = quicvarint.Append(b, uint64(protocol.MaxStreamCount+1)) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "initial_max_streams_uni too large: 1152921504606846977 (maximum 1152921504606846976)", + }, + { + name: "not enough data to read", + data: func() []byte { + b := quicvarint.Append(nil, 0x42) + b = quicvarint.Append(b, 7) + return append(b, []byte("foobar")...) + }(), + perspective: protocol.PerspectiveServer, + expectedErrMsg: "remaining length (6) smaller than parameter length (7)", + }, + { + name: "client sent stateless reset token", + data: func() []byte { + b := quicvarint.Append(nil, uint64(statelessResetTokenParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(16))) + return append(b, make([]byte, 16)...) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "client sent a stateless_reset_token", + }, + { + name: "client sent original destination connection ID", + data: func() []byte { + b := quicvarint.Append(nil, uint64(originalDestinationConnectionIDParameterID)) + b = quicvarint.Append(b, 6) + return append(b, []byte("foobar")...) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "client sent an original_destination_connection_id", + }, + { + name: "huge max ack delay value", + data: func() []byte { + val := uint64(math.MaxUint64) / 5 + b := quicvarint.Append(nil, uint64(maxAckDelayParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(val))) + b = quicvarint.Append(b, val) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "invalid value for max_ack_delay: 3689348814741910323ms (maximum 16383ms)", + }, + { + name: "invalid value for reset_stream_at", + data: func() []byte { + b := quicvarint.Append(nil, uint64(resetStreamAtParameterID)) + b = quicvarint.Append(b, 1) + b = quicvarint.Append(b, 1) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "wrong length for reset_stream_at: 1 (expected empty)", + }, + { + name: "invalid value for legacy reset_stream_at", + data: func() []byte { + b := quicvarint.Append(nil, uint64(legacyResetStreamAtParameterID)) + b = quicvarint.Append(b, 1) + b = quicvarint.Append(b, 1) + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "wrong length for reset_stream_at: 1 (expected empty)", + }, + { + name: "min ack delay is greater than max ack delay", + data: func() []byte { + b := quicvarint.Append(nil, uint64(minAckDelayParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(42001))) + b = quicvarint.Append(b, 42001) // 42001 microseconds + b = quicvarint.Append(b, uint64(maxAckDelayParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(42))) + b = quicvarint.Append(b, 42) // 42 microseconds + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "min_ack_delay (42.001ms) is greater than max_ack_delay (42ms)", + }, + { + name: "huge min ack delay value", + data: func() []byte { + b := quicvarint.Append(nil, uint64(minAckDelayParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(quicvarint.Max))) + b = quicvarint.Append(b, quicvarint.Max) + b = quicvarint.Append(b, uint64(maxAckDelayParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(42))) + b = quicvarint.Append(b, 42) // 42 microseconds + return appendInitialSourceConnectionID(b) + }(), + perspective: protocol.PerspectiveClient, + expectedErrMsg: "min_ack_delay (2562047h47m16.854775807s) is greater than max_ack_delay (42ms)", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var err error + if tt.params != nil { + data := tt.params.Marshal(tt.perspective) + err = (&TransportParameters{}).Unmarshal(data, tt.perspective) + } else { + err = (&TransportParameters{}).Unmarshal(tt.data, tt.perspective) + } + require.Error(t, err) + transportErr, ok := err.(*qerr.TransportError) + require.True(t, ok) + require.Equal(t, qerr.TransportParameterError, transportErr.ErrorCode) + require.Equal(t, tt.expectedErrMsg, transportErr.ErrorMessage) + }) + } +} + +func TestTransportParameterUnknownParameters(t *testing.T) { + // write a known parameter + b := quicvarint.Append(nil, uint64(initialMaxStreamDataBidiLocalParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(0x1337))) + b = quicvarint.Append(b, 0x1337) + // write an unknown parameter + b = quicvarint.Append(b, 0x42) + b = quicvarint.Append(b, 6) + b = append(b, []byte("foobar")...) + // write a known parameter + b = quicvarint.Append(b, uint64(initialMaxStreamDataBidiRemoteParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(0x42))) + b = quicvarint.Append(b, 0x42) + b = appendInitialSourceConnectionID(b) + p := &TransportParameters{} + err := p.Unmarshal(b, protocol.PerspectiveClient) + require.NoError(t, err) + require.Equal(t, protocol.ByteCount(0x1337), p.InitialMaxStreamDataBidiLocal) + require.Equal(t, protocol.ByteCount(0x42), p.InitialMaxStreamDataBidiRemote) +} + +func TestSessionTicketTransportParameterRejectsUnknownParameter(t *testing.T) { + b := (&TransportParameters{ + ActiveConnectionIDLimit: 2, + MaxDatagramFrameSize: protocol.InvalidByteCount, + }).MarshalForSessionTicket(nil) + b = quicvarint.Append(b, 0x42) + b = quicvarint.Append(b, 6) + b = append(b, []byte("foobar")...) + + var p TransportParameters + err := p.UnmarshalFromSessionTicket(b) + require.EqualError(t, err, "unknown transport parameter 0x42 in session ticket") +} + +func TestTransportParameterRejectsDuplicateParameters(t *testing.T) { + // write first parameter + b := quicvarint.Append(nil, uint64(initialMaxStreamDataBidiLocalParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(0x1337))) + b = quicvarint.Append(b, 0x1337) + // write a second parameter + b = quicvarint.Append(b, uint64(initialMaxStreamDataBidiRemoteParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(0x42))) + b = quicvarint.Append(b, 0x42) + // write first parameter again + b = quicvarint.Append(b, uint64(initialMaxStreamDataBidiLocalParameterID)) + b = quicvarint.Append(b, uint64(quicvarint.Len(0x1337))) + b = quicvarint.Append(b, 0x1337) + b = appendInitialSourceConnectionID(b) + err := (&TransportParameters{}).Unmarshal(b, protocol.PerspectiveClient) + require.Error(t, err) + transportErr, ok := err.(*qerr.TransportError) + require.True(t, ok) + require.Equal(t, qerr.TransportParameterError, transportErr.ErrorCode) + require.Equal(t, fmt.Sprintf("received duplicate transport parameter %#x", initialMaxStreamDataBidiLocalParameterID), transportErr.ErrorMessage) +} + +func TestTransportParameterPreferredAddress(t *testing.T) { + testCases := []struct { + name string + hasIPv4 bool + hasIPv6 bool + }{ + {"IPv4 and IPv6", true, true}, + {"IPv4 only", true, false}, + {"IPv6 only", false, true}, + {"neither IPv4 nor IPv6", false, false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + testTransportParameterPreferredAddress(t, tc.hasIPv4, tc.hasIPv6) + }) + } +} + +func testTransportParameterPreferredAddress(t *testing.T, hasIPv4, hasIPv6 bool) { + addr4 := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 42) + addr6 := netip.AddrPortFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}), 13) + pa := &PreferredAddress{ + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}, + } + if hasIPv4 { + pa.IPv4 = addr4 + } + if hasIPv6 { + pa.IPv6 = addr6 + } + + data := (&TransportParameters{ + PreferredAddress: pa, + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: 2, + }).Marshal(protocol.PerspectiveServer) + p := &TransportParameters{} + require.NoError(t, p.Unmarshal(data, protocol.PerspectiveServer)) + if hasIPv4 { + require.True(t, p.PreferredAddress.IPv4.IsValid()) + require.Equal(t, addr4, p.PreferredAddress.IPv4) + } else { + require.False(t, p.PreferredAddress.IPv4.IsValid()) + } + if hasIPv6 { + require.True(t, p.PreferredAddress.IPv6.IsValid()) + require.Equal(t, addr6, p.PreferredAddress.IPv6) + } else { + require.False(t, p.PreferredAddress.IPv6.IsValid()) + } + require.Equal(t, pa.ConnectionID, p.PreferredAddress.ConnectionID) + require.Equal(t, pa.StatelessResetToken, p.PreferredAddress.StatelessResetToken) +} + +func TestTransportParameterPreferredAddressFromClient(t *testing.T) { + b := quicvarint.Append(nil, uint64(preferredAddressParameterID)) + b = quicvarint.Append(b, 6) + b = append(b, []byte("foobar")...) + p := &TransportParameters{} + err := p.Unmarshal(b, protocol.PerspectiveClient) + require.Error(t, err) + require.IsType(t, &qerr.TransportError{}, err) + transportErr := err.(*qerr.TransportError) + require.Equal(t, qerr.TransportParameterError, transportErr.ErrorCode) + require.Equal(t, "client sent a preferred_address", transportErr.ErrorMessage) +} + +func TestTransportParameterPreferredAddressZeroLengthConnectionID(t *testing.T) { + pa := &PreferredAddress{ + IPv4: netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 42), + IPv6: netip.AddrPortFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}), 13), + ConnectionID: protocol.ParseConnectionID([]byte{}), + StatelessResetToken: protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}, + } + data := (&TransportParameters{ + PreferredAddress: pa, + StatelessResetToken: &protocol.StatelessResetToken{}, + }).Marshal(protocol.PerspectiveServer) + p := &TransportParameters{} + err := p.Unmarshal(data, protocol.PerspectiveServer) + require.Error(t, err) + require.IsType(t, &qerr.TransportError{}, err) + transportErr := err.(*qerr.TransportError) + require.Equal(t, qerr.TransportParameterError, transportErr.ErrorCode) + require.Equal(t, "invalid connection ID length: 0", transportErr.ErrorMessage) +} + +func TestPreferredAddressErrorOnEOF(t *testing.T) { + raw := []byte{ + 127, 0, 0, 1, // IPv4 + 0, 42, // IPv4 Port + 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, // IPv6 + 13, 37, // IPv6 Port, + 4, // conn ID len + 0xde, 0xad, 0xbe, 0xef, + 16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, // stateless reset token + } + for i := 1; i < len(raw); i++ { + b := quicvarint.Append(nil, uint64(preferredAddressParameterID)) + b = append(b, raw[:i]...) + p := &TransportParameters{} + err := p.Unmarshal(b, protocol.PerspectiveServer) + require.Error(t, err) + } +} + +func TestTransportParametersFromSessionTicket(t *testing.T) { + params := &TransportParameters{ + InitialMaxStreamDataBidiLocal: protocol.ByteCount(getRandomValue()), + InitialMaxStreamDataBidiRemote: protocol.ByteCount(getRandomValue()), + InitialMaxStreamDataUni: protocol.ByteCount(getRandomValue()), + InitialMaxData: protocol.ByteCount(getRandomValue()), + MaxBidiStreamNum: protocol.StreamNum(getRandomValueUpTo(uint64(protocol.MaxStreamCount))), + MaxUniStreamNum: protocol.StreamNum(getRandomValueUpTo(uint64(protocol.MaxStreamCount))), + ActiveConnectionIDLimit: 2 + getRandomValueUpTo(quicvarint.Max-2), + MaxDatagramFrameSize: protocol.ByteCount(getRandomValueUpTo(uint64(MaxDatagramSize))), + EnableResetStreamAt: getRandomValue()%2 == 0, + } + require.True(t, params.ValidFor0RTT(params)) + b := params.MarshalForSessionTicket(nil) + var tp TransportParameters + require.NoError(t, tp.UnmarshalFromSessionTicket(b)) + require.Equal(t, params.InitialMaxStreamDataBidiLocal, tp.InitialMaxStreamDataBidiLocal) + require.Equal(t, params.InitialMaxStreamDataBidiRemote, tp.InitialMaxStreamDataBidiRemote) + require.Equal(t, params.InitialMaxStreamDataUni, tp.InitialMaxStreamDataUni) + require.Equal(t, params.InitialMaxData, tp.InitialMaxData) + require.Equal(t, params.MaxBidiStreamNum, tp.MaxBidiStreamNum) + require.Equal(t, params.MaxUniStreamNum, tp.MaxUniStreamNum) + require.Equal(t, params.ActiveConnectionIDLimit, tp.ActiveConnectionIDLimit) + require.Equal(t, params.MaxDatagramFrameSize, tp.MaxDatagramFrameSize) + require.Equal(t, params.EnableResetStreamAt, tp.EnableResetStreamAt) +} + +func TestSessionTicketInvalidTransportParameters(t *testing.T) { + var p TransportParameters + require.Error(t, p.UnmarshalFromSessionTicket([]byte("foobar"))) +} + +func TestSessionTicketLegacyResetStreamAtTransportParameter(t *testing.T) { + b := quicvarint.Append(nil, transportParameterMarshalingVersion) + b = quicvarint.Append(b, uint64(legacyResetStreamAtParameterID)) + b = quicvarint.Append(b, 0) + + var p TransportParameters + require.NoError(t, p.UnmarshalFromSessionTicket(b)) + require.True(t, p.EnableResetStreamAt) +} + +func TestSessionTicketTransportParameterVersionMismatch(t *testing.T) { + var p TransportParameters + data := p.MarshalForSessionTicket(nil) + b := quicvarint.Append(nil, transportParameterMarshalingVersion+1) + b = append(b, data[quicvarint.Len(transportParameterMarshalingVersion):]...) + err := p.UnmarshalFromSessionTicket(b) + require.EqualError(t, err, fmt.Sprintf("unknown transport parameter marshaling version: %d", transportParameterMarshalingVersion+1)) +} + +func TestTransportParametersValidFor0RTT(t *testing.T) { + saved := &TransportParameters{ + InitialMaxStreamDataBidiLocal: 1, + InitialMaxStreamDataBidiRemote: 2, + InitialMaxStreamDataUni: 3, + InitialMaxData: 4, + MaxBidiStreamNum: 5, + MaxUniStreamNum: 6, + ActiveConnectionIDLimit: 7, + MaxDatagramFrameSize: 1000, + EnableResetStreamAt: true, + } + + tests := []struct { + name string + modify func(*TransportParameters) + valid bool + }{ + { + name: "No Changes", + modify: func(p *TransportParameters) {}, + valid: true, + }, + { + name: "ResetStreamAt disabled", + modify: func(p *TransportParameters) { p.EnableResetStreamAt = false }, + valid: false, + }, + { + name: "InitialMaxStreamDataBidiLocal reduced", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiLocal = saved.InitialMaxStreamDataBidiLocal - 1 + }, + valid: false, + }, + { + name: "InitialMaxStreamDataBidiLocal increased", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiLocal = saved.InitialMaxStreamDataBidiLocal + 1 + }, + valid: true, + }, + { + name: "InitialMaxStreamDataBidiRemote reduced", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiRemote = saved.InitialMaxStreamDataBidiRemote - 1 + }, + valid: false, + }, + { + name: "InitialMaxStreamDataBidiRemote increased", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiRemote = saved.InitialMaxStreamDataBidiRemote + 1 + }, + valid: true, + }, + { + name: "InitialMaxStreamDataUni reduced", + modify: func(p *TransportParameters) { p.InitialMaxStreamDataUni = saved.InitialMaxStreamDataUni - 1 }, + valid: false, + }, + { + name: "InitialMaxStreamDataUni increased", + modify: func(p *TransportParameters) { p.InitialMaxStreamDataUni = saved.InitialMaxStreamDataUni + 1 }, + valid: true, + }, + { + name: "InitialMaxData reduced", + modify: func(p *TransportParameters) { p.InitialMaxData = saved.InitialMaxData - 1 }, + valid: false, + }, + { + name: "InitialMaxData increased", + modify: func(p *TransportParameters) { p.InitialMaxData = saved.InitialMaxData + 1 }, + valid: true, + }, + { + name: "MaxBidiStreamNum reduced", + modify: func(p *TransportParameters) { p.MaxBidiStreamNum = saved.MaxBidiStreamNum - 1 }, + valid: false, + }, + { + name: "MaxBidiStreamNum increased", + modify: func(p *TransportParameters) { p.MaxBidiStreamNum = saved.MaxBidiStreamNum + 1 }, + valid: true, + }, + { + name: "MaxUniStreamNum reduced", + modify: func(p *TransportParameters) { p.MaxUniStreamNum = saved.MaxUniStreamNum - 1 }, + valid: false, + }, + { + name: "MaxUniStreamNum increased", + modify: func(p *TransportParameters) { p.MaxUniStreamNum = saved.MaxUniStreamNum + 1 }, + valid: true, + }, + { + name: "ActiveConnectionIDLimit changed", + modify: func(p *TransportParameters) { p.ActiveConnectionIDLimit = 0 }, + valid: false, + }, + { + name: "MaxDatagramFrameSize increased", + modify: func(p *TransportParameters) { p.MaxDatagramFrameSize = saved.MaxDatagramFrameSize + 1 }, + valid: true, + }, + { + name: "MaxDatagramFrameSize reduced", + modify: func(p *TransportParameters) { p.MaxDatagramFrameSize = saved.MaxDatagramFrameSize - 1 }, + valid: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + p := *saved + tt.modify(&p) + require.Equal(t, tt.valid, p.ValidFor0RTT(saved)) + }) + } + t.Run("ResetStreamAt enabled", func(t *testing.T) { + p := *saved + withoutResetStreamAt := *saved + withoutResetStreamAt.EnableResetStreamAt = false + require.True(t, p.ValidFor0RTT(&withoutResetStreamAt)) + }) +} + +func TestTransportParametersValidAfter0RTT(t *testing.T) { + saved := &TransportParameters{ + InitialMaxStreamDataBidiLocal: 1, + InitialMaxStreamDataBidiRemote: 2, + InitialMaxStreamDataUni: 3, + InitialMaxData: 4, + MaxBidiStreamNum: 5, + MaxUniStreamNum: 6, + ActiveConnectionIDLimit: 7, + MaxDatagramFrameSize: 1000, + EnableResetStreamAt: true, + } + + tests := []struct { + name string + modify func(*TransportParameters) + reject bool + }{ + { + name: "no changes", + modify: func(p *TransportParameters) {}, + reject: false, + }, + { + name: "ResetStreamAt disabled", + modify: func(p *TransportParameters) { p.EnableResetStreamAt = false }, + reject: true, + }, + { + name: "InitialMaxStreamDataBidiLocal reduced", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiLocal = saved.InitialMaxStreamDataBidiLocal - 1 + }, + reject: true, + }, + { + name: "InitialMaxStreamDataBidiLocal increased", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiLocal = saved.InitialMaxStreamDataBidiLocal + 1 + }, + reject: false, + }, + { + name: "InitialMaxStreamDataBidiRemote reduced", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiRemote = saved.InitialMaxStreamDataBidiRemote - 1 + }, + reject: true, + }, + { + name: "InitialMaxStreamDataBidiRemote increased", + modify: func(p *TransportParameters) { + p.InitialMaxStreamDataBidiRemote = saved.InitialMaxStreamDataBidiRemote + 1 + }, + reject: false, + }, + { + name: "InitialMaxStreamDataUni reduced", + modify: func(p *TransportParameters) { p.InitialMaxStreamDataUni = saved.InitialMaxStreamDataUni - 1 }, + reject: true, + }, + { + name: "InitialMaxStreamDataUni increased", + modify: func(p *TransportParameters) { p.InitialMaxStreamDataUni = saved.InitialMaxStreamDataUni + 1 }, + reject: false, + }, + { + name: "InitialMaxData reduced", + modify: func(p *TransportParameters) { p.InitialMaxData = saved.InitialMaxData - 1 }, + reject: true, + }, + { + name: "InitialMaxData increased", + modify: func(p *TransportParameters) { p.InitialMaxData = saved.InitialMaxData + 1 }, + reject: false, + }, + { + name: "MaxBidiStreamNum reduced", + modify: func(p *TransportParameters) { p.MaxBidiStreamNum = saved.MaxBidiStreamNum - 1 }, + reject: true, + }, + { + name: "MaxBidiStreamNum increased", + modify: func(p *TransportParameters) { p.MaxBidiStreamNum = saved.MaxBidiStreamNum + 1 }, + reject: false, + }, + { + name: "MaxUniStreamNum reduced", + modify: func(p *TransportParameters) { p.MaxUniStreamNum = saved.MaxUniStreamNum - 1 }, + reject: true, + }, + { + name: "MaxUniStreamNum increased", + modify: func(p *TransportParameters) { p.MaxUniStreamNum = saved.MaxUniStreamNum + 1 }, + reject: false, + }, + { + name: "ActiveConnectionIDLimit reduced", + modify: func(p *TransportParameters) { p.ActiveConnectionIDLimit = saved.ActiveConnectionIDLimit - 1 }, + reject: true, + }, + { + name: "ActiveConnectionIDLimit increased", + modify: func(p *TransportParameters) { p.ActiveConnectionIDLimit = saved.ActiveConnectionIDLimit + 1 }, + reject: false, + }, + { + name: "MaxDatagramFrameSize reduced", + modify: func(p *TransportParameters) { p.MaxDatagramFrameSize = saved.MaxDatagramFrameSize - 1 }, + reject: true, + }, + { + name: "MaxDatagramFrameSize increased", + modify: func(p *TransportParameters) { p.MaxDatagramFrameSize = saved.MaxDatagramFrameSize + 1 }, + reject: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + p := *saved + tt.modify(&p) + if tt.reject { + require.False(t, p.ValidForUpdate(saved)) + } else { + require.True(t, p.ValidForUpdate(saved)) + } + }) + } + t.Run("ResetStreamAt enabled", func(t *testing.T) { + p := *saved + withoutResetStreamAt := *saved + withoutResetStreamAt.EnableResetStreamAt = false + require.True(t, p.ValidForUpdate(&withoutResetStreamAt)) + }) +} + +func BenchmarkTransportParameters(b *testing.B) { + b.Run("without preferred address", func(b *testing.B) { benchmarkTransportParameters(b, false) }) + b.Run("with preferred address", func(b *testing.B) { benchmarkTransportParameters(b, true) }) +} + +func benchmarkTransportParameters(b *testing.B, withPreferredAddress bool) { + b.ReportAllocs() + + var token protocol.StatelessResetToken + rand.Read(token[:]) + rcid := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}) + params := &TransportParameters{ + InitialMaxStreamDataBidiLocal: protocol.ByteCount(getRandomValue()), + InitialMaxStreamDataBidiRemote: protocol.ByteCount(getRandomValue()), + InitialMaxStreamDataUni: protocol.ByteCount(getRandomValue()), + InitialMaxData: protocol.ByteCount(getRandomValue()), + MaxIdleTimeout: 0xcafe * time.Second, + MaxBidiStreamNum: protocol.StreamNum(getRandomValueUpTo(uint64(protocol.MaxStreamCount))), + MaxUniStreamNum: protocol.StreamNum(getRandomValueUpTo(uint64(protocol.MaxStreamCount))), + DisableActiveMigration: true, + StatelessResetToken: &token, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + RetrySourceConnectionID: &rcid, + AckDelayExponent: 13, + MaxAckDelay: 42 * time.Millisecond, + ActiveConnectionIDLimit: 2 + getRandomValueUpTo(quicvarint.Max-2), + MaxDatagramFrameSize: protocol.ByteCount(getRandomValue()), + } + var token2 protocol.StatelessResetToken + rand.Read(token2[:]) + if withPreferredAddress { + var ip4 [4]byte + var ip6 [16]byte + rand.Read(ip4[:]) + rand.Read(ip6[:]) + params.PreferredAddress = &PreferredAddress{ + IPv4: netip.AddrPortFrom(netip.AddrFrom4(ip4), 1234), + IPv6: netip.AddrPortFrom(netip.AddrFrom16(ip6), 4321), + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: token2, + } + } + data := params.Marshal(protocol.PerspectiveServer) + + var p TransportParameters + for b.Loop() { + if err := p.Unmarshal(data, protocol.PerspectiveServer); err != nil { + b.Fatal(err) + } + // check a few fields + if p.DisableActiveMigration != params.DisableActiveMigration || + p.InitialMaxStreamDataBidiLocal != params.InitialMaxStreamDataBidiLocal || + *p.StatelessResetToken != *params.StatelessResetToken || + p.AckDelayExponent != params.AckDelayExponent { + b.Fatalf("params mismatch: %v vs %v", p, params) + } + if withPreferredAddress && *p.PreferredAddress != *params.PreferredAddress { + b.Fatalf("preferred address mismatch: %v vs %v", p.PreferredAddress, params.PreferredAddress) + } + } +} + +func FuzzTransportParameters(f *testing.F) { + corpus := ossfuzzseeds.New(f) + + savedParams := (&TransportParameters{ + InitialMaxStreamDataBidiLocal: 1234, + InitialMaxStreamDataBidiRemote: 2345, + InitialMaxStreamDataUni: 3456, + InitialMaxData: 4567, + MaxBidiStreamNum: 1337, + MaxUniStreamNum: 7331, + ActiveConnectionIDLimit: 7, + MaxDatagramFrameSize: protocol.InvalidByteCount, + }).MarshalForSessionTicket(nil) + zeroRTTParams := (&TransportParameters{ + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: 7, + MaxDatagramFrameSize: 1200, + }).Marshal(protocol.PerspectiveServer) + + rcid := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}) + minAckDelay := 42 * time.Millisecond + for _, seed := range []struct { + Data []byte + SavedData []byte + }{ + {(&TransportParameters{ + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: 2, + }).Marshal(protocol.PerspectiveServer), savedParams}, + {zeroRTTParams, savedParams}, + {(&TransportParameters{ + ActiveConnectionIDLimit: 2, + RetrySourceConnectionID: &rcid, + }).Marshal(protocol.PerspectiveClient), savedParams}, + {(&TransportParameters{ + EnableResetStreamAt: true, + RetrySourceConnectionID: &rcid, + MinAckDelay: &minAckDelay, + }).Marshal(protocol.PerspectiveClient), savedParams}, + // session ticket + {savedParams, savedParams}, + // with preferred address + {(&TransportParameters{ + StatelessResetToken: &protocol.StatelessResetToken{}, + ActiveConnectionIDLimit: 2, + PreferredAddress: &PreferredAddress{ + IPv4: netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 42), + IPv6: netip.AddrPortFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}), 13), + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1}, + }, + }).Marshal(protocol.PerspectiveServer), savedParams}, + } { + corpus.Add(seed.Data, seed.SavedData) + } + + f.Fuzz(func(t *testing.T, data, savedData []byte) { + fuzzTransportParameters(t, data, protocol.PerspectiveClient) + fuzzTransportParameters(t, data, protocol.PerspectiveServer) + fuzzTransportParametersSessionTicket(t, data) + fuzzTransportParameters0RTT(t, data, savedData) + }) +} + +func fuzzTransportParameters(t *testing.T, data []byte, sentBy protocol.Perspective) { + t.Helper() + + tp := &TransportParameters{} + if err := tp.Unmarshal(data, sentBy); err != nil { + return + } + _ = tp.String() + checkTransportParameterInvariants(t, tp, sentBy) + + tp2 := &TransportParameters{} + if err := tp2.Unmarshal(tp.Marshal(sentBy), sentBy); err != nil { + t.Fatalf("error unmarshaling re-marshaled transport parameters: %s", err) + } + checkTransportParameterInvariants(t, tp2, sentBy) +} + +func fuzzTransportParametersSessionTicket(t *testing.T, data []byte) { + t.Helper() + + tp := &TransportParameters{} + if err := tp.UnmarshalFromSessionTicket(data); err != nil { + return + } + _ = tp.String() + b := tp.MarshalForSessionTicket(nil) + tp2 := &TransportParameters{} + if err := tp2.UnmarshalFromSessionTicket(b); err != nil { + t.Fatalf("error unmarshaling re-marshaled session ticket transport parameters: %s", err) + } +} + +func fuzzTransportParameters0RTT(t *testing.T, data, savedData []byte) { + t.Helper() + + tp := &TransportParameters{} + if err := tp.Unmarshal(data, protocol.PerspectiveServer); err != nil { + return + } + saved := &TransportParameters{} + if err := saved.UnmarshalFromSessionTicket(savedData); err != nil { + return + } + _ = tp.ValidFor0RTT(saved) + _ = tp.ValidForUpdate(saved) +} + +func checkTransportParameterInvariants(t *testing.T, tp *TransportParameters, sentBy protocol.Perspective) { + t.Helper() + + if sentBy == protocol.PerspectiveClient && tp.StatelessResetToken != nil { + t.Fatal("client's transport parameters contained stateless reset token") + } + if tp.MaxIdleTimeout < 0 { + t.Fatalf("negative max_idle_timeout: %s", tp.MaxIdleTimeout) + } + if tp.AckDelayExponent > 20 { + t.Fatalf("invalid ack_delay_exponent: %d", tp.AckDelayExponent) + } + if tp.MaxUDPPayloadSize < 1200 { + t.Fatalf("invalid max_udp_payload_size: %d", tp.MaxUDPPayloadSize) + } + if tp.ActiveConnectionIDLimit < 2 { + t.Fatalf("invalid active_connection_id_limit: %d", tp.ActiveConnectionIDLimit) + } + if tp.OriginalDestinationConnectionID.Len() > 20 { + t.Fatalf("invalid original_destination_connection_id length: %s", tp.OriginalDestinationConnectionID) + } + if tp.InitialSourceConnectionID.Len() > 20 { + t.Fatalf("invalid initial_source_connection_id length: %s", tp.InitialSourceConnectionID) + } + if tp.RetrySourceConnectionID != nil && tp.RetrySourceConnectionID.Len() > 20 { + t.Fatalf("invalid retry_source_connection_id length: %s", tp.RetrySourceConnectionID) + } + if tp.PreferredAddress != nil && tp.PreferredAddress.ConnectionID.Len() > 20 { + t.Fatalf("invalid preferred_address connection ID length: %s", tp.PreferredAddress.ConnectionID) + } + if tp.MinAckDelay != nil { + if *tp.MinAckDelay < 0 { + t.Fatalf("negative min_ack_delay: %s", *tp.MinAckDelay) + } + if *tp.MinAckDelay > tp.MaxAckDelay { + t.Fatalf("min_ack_delay (%s) is greater than max_ack_delay (%s)", *tp.MinAckDelay, tp.MaxAckDelay) + } + } +} diff --git a/third_party/quic-go/internal/wire/transport_parameters.go b/third_party/quic-go/internal/wire/transport_parameters.go new file mode 100644 index 0000000..a493e89 --- /dev/null +++ b/third_party/quic-go/internal/wire/transport_parameters.go @@ -0,0 +1,601 @@ +package wire + +import ( + "crypto/rand" + "encoding/binary" + "errors" + "fmt" + "io" + "math" + "net/netip" + "slices" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/quicvarint" +) + +// AdditionalTransportParametersClient are additional transport parameters that will be added +// to the client's transport parameters. +// This is not intended for production use, but _only_ to increase the size of the ClientHello beyond +// the usual size of less than 1 MTU. +var AdditionalTransportParametersClient map[uint64][]byte + +const transportParameterMarshalingVersion = 1 + +type transportParameterID uint64 + +const ( + originalDestinationConnectionIDParameterID transportParameterID = 0x0 + maxIdleTimeoutParameterID transportParameterID = 0x1 + statelessResetTokenParameterID transportParameterID = 0x2 + maxUDPPayloadSizeParameterID transportParameterID = 0x3 + initialMaxDataParameterID transportParameterID = 0x4 + initialMaxStreamDataBidiLocalParameterID transportParameterID = 0x5 + initialMaxStreamDataBidiRemoteParameterID transportParameterID = 0x6 + initialMaxStreamDataUniParameterID transportParameterID = 0x7 + initialMaxStreamsBidiParameterID transportParameterID = 0x8 + initialMaxStreamsUniParameterID transportParameterID = 0x9 + ackDelayExponentParameterID transportParameterID = 0xa + maxAckDelayParameterID transportParameterID = 0xb + disableActiveMigrationParameterID transportParameterID = 0xc + preferredAddressParameterID transportParameterID = 0xd + activeConnectionIDLimitParameterID transportParameterID = 0xe + initialSourceConnectionIDParameterID transportParameterID = 0xf + retrySourceConnectionIDParameterID transportParameterID = 0x10 + // RFC 9221 + maxDatagramFrameSizeParameterID transportParameterID = 0x20 + // https://datatracker.ietf.org/doc/draft-ietf-quic-reliable-stream-reset/09/ + resetStreamAtParameterID transportParameterID = 0x1d + // https://datatracker.ietf.org/doc/draft-ietf-quic-reliable-stream-reset/07/ + // When removing support for this codepoint, increment transportParameterMarshalingVersion + // to prevent 0-RTT resumption with tickets that remember it. + legacyResetStreamAtParameterID transportParameterID = 0x17f7586d2cb571 + // https://datatracker.ietf.org/doc/draft-ietf-quic-ack-frequency/11/ + minAckDelayParameterID transportParameterID = 0xff04de1b +) + +// PreferredAddress is the value encoding in the preferred_address transport parameter +type PreferredAddress struct { + IPv4, IPv6 netip.AddrPort + ConnectionID protocol.ConnectionID + StatelessResetToken protocol.StatelessResetToken +} + +// TransportParameters are parameters sent to the peer during the handshake +type TransportParameters struct { + InitialMaxStreamDataBidiLocal protocol.ByteCount + InitialMaxStreamDataBidiRemote protocol.ByteCount + InitialMaxStreamDataUni protocol.ByteCount + InitialMaxData protocol.ByteCount + + MaxAckDelay time.Duration + AckDelayExponent uint8 + + DisableActiveMigration bool + + MaxUDPPayloadSize protocol.ByteCount + + MaxUniStreamNum protocol.StreamNum + MaxBidiStreamNum protocol.StreamNum + + MaxIdleTimeout time.Duration + + PreferredAddress *PreferredAddress + + OriginalDestinationConnectionID protocol.ConnectionID + InitialSourceConnectionID protocol.ConnectionID + RetrySourceConnectionID *protocol.ConnectionID // use a pointer here to distinguish zero-length connection IDs from missing transport parameters + + StatelessResetToken *protocol.StatelessResetToken + ActiveConnectionIDLimit uint64 + + MaxDatagramFrameSize protocol.ByteCount // RFC 9221 + EnableResetStreamAt bool // https://datatracker.ietf.org/doc/draft-ietf-quic-reliable-stream-reset/09/ + MinAckDelay *time.Duration + + // ChromeFingerprint makes Marshal encode the parameters the way Google + // Chrome does. Client side only; see marshalChrome. + ChromeFingerprint bool +} + +// Unmarshal the transport parameters +func (p *TransportParameters) Unmarshal(data []byte, sentBy protocol.Perspective) error { + if err := p.unmarshal(data, sentBy, false); err != nil { + return &qerr.TransportError{ + ErrorCode: qerr.TransportParameterError, + ErrorMessage: err.Error(), + } + } + return nil +} + +func (p *TransportParameters) unmarshal(b []byte, sentBy protocol.Perspective, fromSessionTicket bool) error { + // needed to check that every parameter is only sent at most once + parameterIDs := make([]transportParameterID, 0, 32) + + var ( + readOriginalDestinationConnectionID bool + readInitialSourceConnectionID bool + ) + + p.AckDelayExponent = protocol.DefaultAckDelayExponent + p.MaxAckDelay = protocol.DefaultMaxAckDelay + p.MaxDatagramFrameSize = protocol.InvalidByteCount + p.ActiveConnectionIDLimit = protocol.DefaultActiveConnectionIDLimit + + for len(b) > 0 { + paramIDInt, l, err := quicvarint.Parse(b) + if err != nil { + return err + } + paramID := transportParameterID(paramIDInt) + b = b[l:] + paramLen, l, err := quicvarint.Parse(b) + if err != nil { + return err + } + b = b[l:] + if uint64(len(b)) < paramLen { + return fmt.Errorf("remaining length (%d) smaller than parameter length (%d)", len(b), paramLen) + } + parameterIDs = append(parameterIDs, paramID) + switch paramID { + case maxIdleTimeoutParameterID, + maxUDPPayloadSizeParameterID, + initialMaxDataParameterID, + initialMaxStreamDataBidiLocalParameterID, + initialMaxStreamDataBidiRemoteParameterID, + initialMaxStreamDataUniParameterID, + initialMaxStreamsBidiParameterID, + initialMaxStreamsUniParameterID, + maxAckDelayParameterID, + maxDatagramFrameSizeParameterID, + ackDelayExponentParameterID, + activeConnectionIDLimitParameterID, + minAckDelayParameterID: + if err := p.readNumericTransportParameter(b, paramID, int(paramLen)); err != nil { + return err + } + b = b[paramLen:] + case preferredAddressParameterID: + if sentBy == protocol.PerspectiveClient { + return errors.New("client sent a preferred_address") + } + if err := p.readPreferredAddress(b, int(paramLen)); err != nil { + return err + } + b = b[paramLen:] + case disableActiveMigrationParameterID: + if paramLen != 0 { + return fmt.Errorf("wrong length for disable_active_migration: %d (expected empty)", paramLen) + } + p.DisableActiveMigration = true + case statelessResetTokenParameterID: + if sentBy == protocol.PerspectiveClient { + return errors.New("client sent a stateless_reset_token") + } + if paramLen != 16 { + return fmt.Errorf("wrong length for stateless_reset_token: %d (expected 16)", paramLen) + } + var token protocol.StatelessResetToken + if len(b) < len(token) { + return io.EOF + } + copy(token[:], b) + b = b[len(token):] + p.StatelessResetToken = &token + case originalDestinationConnectionIDParameterID: + if sentBy == protocol.PerspectiveClient { + return errors.New("client sent an original_destination_connection_id") + } + if paramLen > protocol.MaxConnIDLen { + return protocol.ErrInvalidConnectionIDLen + } + p.OriginalDestinationConnectionID = protocol.ParseConnectionID(b[:paramLen]) + b = b[paramLen:] + readOriginalDestinationConnectionID = true + case initialSourceConnectionIDParameterID: + if paramLen > protocol.MaxConnIDLen { + return protocol.ErrInvalidConnectionIDLen + } + p.InitialSourceConnectionID = protocol.ParseConnectionID(b[:paramLen]) + b = b[paramLen:] + readInitialSourceConnectionID = true + case retrySourceConnectionIDParameterID: + if sentBy == protocol.PerspectiveClient { + return errors.New("client sent a retry_source_connection_id") + } + if paramLen > protocol.MaxConnIDLen { + return protocol.ErrInvalidConnectionIDLen + } + connID := protocol.ParseConnectionID(b[:paramLen]) + b = b[paramLen:] + p.RetrySourceConnectionID = &connID + case resetStreamAtParameterID, legacyResetStreamAtParameterID: + if paramLen != 0 { + return fmt.Errorf("wrong length for reset_stream_at: %d (expected empty)", paramLen) + } + p.EnableResetStreamAt = true + default: + if fromSessionTicket { + // A ticket might contain a parameter for an extension supported by an older + // version of this endpoint. If we can't parse it, don't resume with it. + return fmt.Errorf("unknown transport parameter %#x in session ticket", paramID) + } + b = b[paramLen:] + } + } + + // min_ack_delay must be less or equal to max_ack_delay + if p.MinAckDelay != nil && *p.MinAckDelay > p.MaxAckDelay { + return fmt.Errorf("min_ack_delay (%s) is greater than max_ack_delay (%s)", *p.MinAckDelay, p.MaxAckDelay) + } + if !fromSessionTicket { + if sentBy == protocol.PerspectiveServer && !readOriginalDestinationConnectionID { + return errors.New("missing original_destination_connection_id") + } + if p.MaxUDPPayloadSize == 0 { + p.MaxUDPPayloadSize = protocol.MaxByteCount + } + if !readInitialSourceConnectionID { + return errors.New("missing initial_source_connection_id") + } + } + + // check that every transport parameter was sent at most once + slices.Sort(parameterIDs) + for i := range len(parameterIDs) - 1 { + if parameterIDs[i] == parameterIDs[i+1] { + return fmt.Errorf("received duplicate transport parameter %#x", parameterIDs[i]) + } + } + + return nil +} + +func (p *TransportParameters) readPreferredAddress(b []byte, expectedLen int) error { + remainingLen := len(b) + pa := &PreferredAddress{} + if len(b) < 4+2+16+2+1 { + return io.EOF + } + var ipv4 [4]byte + copy(ipv4[:], b[:4]) + port4 := binary.BigEndian.Uint16(b[4:]) + b = b[4+2:] + if port4 != 0 && ipv4 != [4]byte{} { + pa.IPv4 = netip.AddrPortFrom(netip.AddrFrom4(ipv4), port4) + } + var ipv6 [16]byte + copy(ipv6[:], b[:16]) + port6 := binary.BigEndian.Uint16(b[16:]) + if port6 != 0 && ipv6 != [16]byte{} { + pa.IPv6 = netip.AddrPortFrom(netip.AddrFrom16(ipv6), port6) + } + b = b[16+2:] + connIDLen := int(b[0]) + b = b[1:] + if connIDLen == 0 || connIDLen > protocol.MaxConnIDLen { + return fmt.Errorf("invalid connection ID length: %d", connIDLen) + } + if len(b) < connIDLen+len(pa.StatelessResetToken) { + return io.EOF + } + pa.ConnectionID = protocol.ParseConnectionID(b[:connIDLen]) + b = b[connIDLen:] + copy(pa.StatelessResetToken[:], b) + b = b[len(pa.StatelessResetToken):] + if bytesRead := remainingLen - len(b); bytesRead != expectedLen { + return fmt.Errorf("expected preferred_address to be %d long, read %d bytes", expectedLen, bytesRead) + } + p.PreferredAddress = pa + return nil +} + +func (p *TransportParameters) readNumericTransportParameter(b []byte, paramID transportParameterID, expectedLen int) error { + val, l, err := quicvarint.Parse(b) + if err != nil { + return fmt.Errorf("error while reading transport parameter %d: %s", paramID, err) + } + if l != expectedLen { + return fmt.Errorf("inconsistent transport parameter length for transport parameter %#x", paramID) + } + //nolint:exhaustive // This only covers the numeric transport parameters. + switch paramID { + case initialMaxStreamDataBidiLocalParameterID: + p.InitialMaxStreamDataBidiLocal = protocol.ByteCount(val) + case initialMaxStreamDataBidiRemoteParameterID: + p.InitialMaxStreamDataBidiRemote = protocol.ByteCount(val) + case initialMaxStreamDataUniParameterID: + p.InitialMaxStreamDataUni = protocol.ByteCount(val) + case initialMaxDataParameterID: + p.InitialMaxData = protocol.ByteCount(val) + case initialMaxStreamsBidiParameterID: + p.MaxBidiStreamNum = protocol.StreamNum(val) + if p.MaxBidiStreamNum > protocol.MaxStreamCount { + return fmt.Errorf("initial_max_streams_bidi too large: %d (maximum %d)", p.MaxBidiStreamNum, protocol.MaxStreamCount) + } + case initialMaxStreamsUniParameterID: + p.MaxUniStreamNum = protocol.StreamNum(val) + if p.MaxUniStreamNum > protocol.MaxStreamCount { + return fmt.Errorf("initial_max_streams_uni too large: %d (maximum %d)", p.MaxUniStreamNum, protocol.MaxStreamCount) + } + case maxIdleTimeoutParameterID: + p.MaxIdleTimeout = max(protocol.MinRemoteIdleTimeout, time.Duration(val)*time.Millisecond) + case maxUDPPayloadSizeParameterID: + if val < 1200 { + return fmt.Errorf("invalid value for max_udp_payload_size: %d (minimum 1200)", val) + } + p.MaxUDPPayloadSize = protocol.ByteCount(val) + case ackDelayExponentParameterID: + if val > protocol.MaxAckDelayExponent { + return fmt.Errorf("invalid value for ack_delay_exponent: %d (maximum %d)", val, protocol.MaxAckDelayExponent) + } + p.AckDelayExponent = uint8(val) + case maxAckDelayParameterID: + if val > uint64(protocol.MaxMaxAckDelay/time.Millisecond) { + return fmt.Errorf("invalid value for max_ack_delay: %dms (maximum %dms)", val, protocol.MaxMaxAckDelay/time.Millisecond) + } + p.MaxAckDelay = time.Duration(val) * time.Millisecond + case activeConnectionIDLimitParameterID: + if val < 2 { + return fmt.Errorf("invalid value for active_connection_id_limit: %d (minimum 2)", val) + } + p.ActiveConnectionIDLimit = val + case maxDatagramFrameSizeParameterID: + p.MaxDatagramFrameSize = protocol.ByteCount(val) + case minAckDelayParameterID: + mad := time.Duration(val) * time.Microsecond + if mad < 0 { + mad = math.MaxInt64 + } + p.MinAckDelay = &mad + default: + return fmt.Errorf("TransportParameter BUG: transport parameter %d not found", paramID) + } + return nil +} + +// Marshal the transport parameters +func (p *TransportParameters) Marshal(pers protocol.Perspective) []byte { + if p.ChromeFingerprint && pers == protocol.PerspectiveClient { + return p.marshalChrome() + } + + // Typical Transport Parameters consume around 110 bytes, depending on the exact values, + // especially the lengths of the Connection IDs. + // Allocate 256 bytes, so we won't have to grow the slice in any case. + b := make([]byte, 0, 256) + + // add a greased value + random := make([]byte, 18) + rand.Read(random) + b = quicvarint.Append(b, 27+31*uint64(random[0])) + length := random[1] % 16 + b = quicvarint.Append(b, uint64(length)) + b = append(b, random[2:2+length]...) + + // initial_max_stream_data_bidi_local + b = p.marshalVarintParam(b, initialMaxStreamDataBidiLocalParameterID, uint64(p.InitialMaxStreamDataBidiLocal)) + // initial_max_stream_data_bidi_remote + b = p.marshalVarintParam(b, initialMaxStreamDataBidiRemoteParameterID, uint64(p.InitialMaxStreamDataBidiRemote)) + // initial_max_stream_data_uni + b = p.marshalVarintParam(b, initialMaxStreamDataUniParameterID, uint64(p.InitialMaxStreamDataUni)) + // initial_max_data + b = p.marshalVarintParam(b, initialMaxDataParameterID, uint64(p.InitialMaxData)) + // initial_max_bidi_streams + b = p.marshalVarintParam(b, initialMaxStreamsBidiParameterID, uint64(p.MaxBidiStreamNum)) + // initial_max_uni_streams + b = p.marshalVarintParam(b, initialMaxStreamsUniParameterID, uint64(p.MaxUniStreamNum)) + // idle_timeout + b = p.marshalVarintParam(b, maxIdleTimeoutParameterID, uint64(p.MaxIdleTimeout/time.Millisecond)) + // max_udp_payload_size + if p.MaxUDPPayloadSize > 0 { + b = p.marshalVarintParam(b, maxUDPPayloadSizeParameterID, uint64(p.MaxUDPPayloadSize)) + } + // max_ack_delay + // Only send it if is different from the default value. + if p.MaxAckDelay != protocol.DefaultMaxAckDelay { + b = p.marshalVarintParam(b, maxAckDelayParameterID, uint64(p.MaxAckDelay/time.Millisecond)) + } + // ack_delay_exponent + // Only send it if is different from the default value. + if p.AckDelayExponent != protocol.DefaultAckDelayExponent { + b = p.marshalVarintParam(b, ackDelayExponentParameterID, uint64(p.AckDelayExponent)) + } + // disable_active_migration + if p.DisableActiveMigration { + b = quicvarint.Append(b, uint64(disableActiveMigrationParameterID)) + b = quicvarint.Append(b, 0) + } + if pers == protocol.PerspectiveServer { + // stateless_reset_token + if p.StatelessResetToken != nil { + b = quicvarint.Append(b, uint64(statelessResetTokenParameterID)) + b = quicvarint.Append(b, 16) + b = append(b, p.StatelessResetToken[:]...) + } + // original_destination_connection_id + b = quicvarint.Append(b, uint64(originalDestinationConnectionIDParameterID)) + b = quicvarint.Append(b, uint64(p.OriginalDestinationConnectionID.Len())) + b = append(b, p.OriginalDestinationConnectionID.Bytes()...) + // preferred_address + if p.PreferredAddress != nil { + b = quicvarint.Append(b, uint64(preferredAddressParameterID)) + b = quicvarint.Append(b, 4+2+16+2+1+uint64(p.PreferredAddress.ConnectionID.Len())+16) + if p.PreferredAddress.IPv4.IsValid() { + ipv4 := p.PreferredAddress.IPv4.Addr().As4() + b = append(b, ipv4[:]...) + b = binary.BigEndian.AppendUint16(b, p.PreferredAddress.IPv4.Port()) + } else { + b = append(b, make([]byte, 6)...) + } + if p.PreferredAddress.IPv6.IsValid() { + ipv6 := p.PreferredAddress.IPv6.Addr().As16() + b = append(b, ipv6[:]...) + b = binary.BigEndian.AppendUint16(b, p.PreferredAddress.IPv6.Port()) + } else { + b = append(b, make([]byte, 18)...) + } + b = append(b, uint8(p.PreferredAddress.ConnectionID.Len())) + b = append(b, p.PreferredAddress.ConnectionID.Bytes()...) + b = append(b, p.PreferredAddress.StatelessResetToken[:]...) + } + } + // active_connection_id_limit + if p.ActiveConnectionIDLimit != protocol.DefaultActiveConnectionIDLimit { + b = p.marshalVarintParam(b, activeConnectionIDLimitParameterID, p.ActiveConnectionIDLimit) + } + // initial_source_connection_id + b = quicvarint.Append(b, uint64(initialSourceConnectionIDParameterID)) + b = quicvarint.Append(b, uint64(p.InitialSourceConnectionID.Len())) + b = append(b, p.InitialSourceConnectionID.Bytes()...) + // retry_source_connection_id + if pers == protocol.PerspectiveServer && p.RetrySourceConnectionID != nil { + b = quicvarint.Append(b, uint64(retrySourceConnectionIDParameterID)) + b = quicvarint.Append(b, uint64(p.RetrySourceConnectionID.Len())) + b = append(b, p.RetrySourceConnectionID.Bytes()...) + } + // QUIC datagrams + if p.MaxDatagramFrameSize != protocol.InvalidByteCount { + b = p.marshalVarintParam(b, maxDatagramFrameSizeParameterID, uint64(p.MaxDatagramFrameSize)) + } + // QUIC Stream Resets with Partial Delivery + if p.EnableResetStreamAt { + b = quicvarint.Append(b, uint64(resetStreamAtParameterID)) + b = quicvarint.Append(b, 0) + } + if p.MinAckDelay != nil { + b = p.marshalVarintParam(b, minAckDelayParameterID, uint64(*p.MinAckDelay/time.Microsecond)) + } + + if pers == protocol.PerspectiveClient && len(AdditionalTransportParametersClient) > 0 { + for k, v := range AdditionalTransportParametersClient { + b = quicvarint.Append(b, k) + b = quicvarint.Append(b, uint64(len(v))) + b = append(b, v...) + } + } + + return b +} + +func (p *TransportParameters) marshalVarintParam(b []byte, id transportParameterID, val uint64) []byte { + b = quicvarint.Append(b, uint64(id)) + b = quicvarint.Append(b, uint64(quicvarint.Len(val))) + return quicvarint.Append(b, val) +} + +// MarshalForSessionTicket marshals the transport parameters we save in the session ticket. +// When sending a 0-RTT enabled TLS session tickets, we need to save the transport parameters. +// The client will remember the transport parameters used in the last session, +// and apply those to the 0-RTT data it sends. +// Saving the transport parameters in the ticket gives the server the option to reject 0-RTT +// if the transport parameters changed. +// Since the session ticket is encrypted, the serialization format is defined by the server. +// For convenience, we use the same format that we also use for sending the transport parameters. +func (p *TransportParameters) MarshalForSessionTicket(b []byte) []byte { + b = quicvarint.Append(b, transportParameterMarshalingVersion) + + // initial_max_stream_data_bidi_local + b = p.marshalVarintParam(b, initialMaxStreamDataBidiLocalParameterID, uint64(p.InitialMaxStreamDataBidiLocal)) + // initial_max_stream_data_bidi_remote + b = p.marshalVarintParam(b, initialMaxStreamDataBidiRemoteParameterID, uint64(p.InitialMaxStreamDataBidiRemote)) + // initial_max_stream_data_uni + b = p.marshalVarintParam(b, initialMaxStreamDataUniParameterID, uint64(p.InitialMaxStreamDataUni)) + // initial_max_data + b = p.marshalVarintParam(b, initialMaxDataParameterID, uint64(p.InitialMaxData)) + // initial_max_bidi_streams + b = p.marshalVarintParam(b, initialMaxStreamsBidiParameterID, uint64(p.MaxBidiStreamNum)) + // initial_max_uni_streams + b = p.marshalVarintParam(b, initialMaxStreamsUniParameterID, uint64(p.MaxUniStreamNum)) + // active_connection_id_limit + b = p.marshalVarintParam(b, activeConnectionIDLimitParameterID, p.ActiveConnectionIDLimit) + // max_datagram_frame_size + if p.MaxDatagramFrameSize != protocol.InvalidByteCount { + b = p.marshalVarintParam(b, maxDatagramFrameSizeParameterID, uint64(p.MaxDatagramFrameSize)) + } + // reset_stream_at + if p.EnableResetStreamAt { + b = quicvarint.Append(b, uint64(resetStreamAtParameterID)) + b = quicvarint.Append(b, 0) + } + return b +} + +// UnmarshalFromSessionTicket unmarshals transport parameters from a session ticket. +func (p *TransportParameters) UnmarshalFromSessionTicket(b []byte) error { + version, l, err := quicvarint.Parse(b) + if err != nil { + return err + } + if version != transportParameterMarshalingVersion { + return fmt.Errorf("unknown transport parameter marshaling version: %d", version) + } + return p.unmarshal(b[l:], protocol.PerspectiveServer, true) +} + +// ValidFor0RTT checks if the transport parameters match those saved in the session ticket. +func (p *TransportParameters) ValidFor0RTT(saved *TransportParameters) bool { + if saved.MaxDatagramFrameSize != protocol.InvalidByteCount && (p.MaxDatagramFrameSize == protocol.InvalidByteCount || p.MaxDatagramFrameSize < saved.MaxDatagramFrameSize) { + return false + } + if saved.EnableResetStreamAt && !p.EnableResetStreamAt { + return false + } + return p.InitialMaxStreamDataBidiLocal >= saved.InitialMaxStreamDataBidiLocal && + p.InitialMaxStreamDataBidiRemote >= saved.InitialMaxStreamDataBidiRemote && + p.InitialMaxStreamDataUni >= saved.InitialMaxStreamDataUni && + p.InitialMaxData >= saved.InitialMaxData && + p.MaxBidiStreamNum >= saved.MaxBidiStreamNum && + p.MaxUniStreamNum >= saved.MaxUniStreamNum && + p.ActiveConnectionIDLimit == saved.ActiveConnectionIDLimit +} + +// ValidForUpdate checks that the new transport parameters don't reduce limits after resuming a 0-RTT connection. +// It is only used on the client side. +func (p *TransportParameters) ValidForUpdate(saved *TransportParameters) bool { + if saved.MaxDatagramFrameSize != protocol.InvalidByteCount && (p.MaxDatagramFrameSize == protocol.InvalidByteCount || p.MaxDatagramFrameSize < saved.MaxDatagramFrameSize) { + return false + } + if saved.EnableResetStreamAt && !p.EnableResetStreamAt { + return false + } + return p.ActiveConnectionIDLimit >= saved.ActiveConnectionIDLimit && + p.InitialMaxData >= saved.InitialMaxData && + p.InitialMaxStreamDataBidiLocal >= saved.InitialMaxStreamDataBidiLocal && + p.InitialMaxStreamDataBidiRemote >= saved.InitialMaxStreamDataBidiRemote && + p.InitialMaxStreamDataUni >= saved.InitialMaxStreamDataUni && + p.MaxBidiStreamNum >= saved.MaxBidiStreamNum && + p.MaxUniStreamNum >= saved.MaxUniStreamNum +} + +// String returns a string representation, intended for logging. +func (p *TransportParameters) String() string { + logString := "&wire.TransportParameters{OriginalDestinationConnectionID: %s, InitialSourceConnectionID: %s, " + logParams := []any{p.OriginalDestinationConnectionID, p.InitialSourceConnectionID} + if p.RetrySourceConnectionID != nil { + logString += "RetrySourceConnectionID: %s, " + logParams = append(logParams, p.RetrySourceConnectionID) + } + logString += "InitialMaxStreamDataBidiLocal: %d, InitialMaxStreamDataBidiRemote: %d, InitialMaxStreamDataUni: %d, InitialMaxData: %d, MaxBidiStreamNum: %d, MaxUniStreamNum: %d, MaxIdleTimeout: %s, AckDelayExponent: %d, MaxAckDelay: %s, ActiveConnectionIDLimit: %d" + logParams = append(logParams, []any{p.InitialMaxStreamDataBidiLocal, p.InitialMaxStreamDataBidiRemote, p.InitialMaxStreamDataUni, p.InitialMaxData, p.MaxBidiStreamNum, p.MaxUniStreamNum, p.MaxIdleTimeout, p.AckDelayExponent, p.MaxAckDelay, p.ActiveConnectionIDLimit}...) + if p.StatelessResetToken != nil { // the client never sends a stateless reset token + logString += ", StatelessResetToken: %#x" + logParams = append(logParams, *p.StatelessResetToken) + } + if p.MaxDatagramFrameSize != protocol.InvalidByteCount { + logString += ", MaxDatagramFrameSize: %d" + logParams = append(logParams, p.MaxDatagramFrameSize) + } + logString += ", EnableResetStreamAt: %t" + logParams = append(logParams, p.EnableResetStreamAt) + if p.MinAckDelay != nil { + logString += ", MinAckDelay: %s" + logParams = append(logParams, *p.MinAckDelay) + } + logString += "}" + return fmt.Sprintf(logString, logParams...) +} diff --git a/third_party/quic-go/internal/wire/transport_parameters_chrome.go b/third_party/quic-go/internal/wire/transport_parameters_chrome.go new file mode 100644 index 0000000..5b56cb9 --- /dev/null +++ b/third_party/quic-go/internal/wire/transport_parameters_chrome.go @@ -0,0 +1,156 @@ +package wire + +import ( + "crypto/rand" + "encoding/binary" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" +) + +// Transport parameters sent by the parroted client that aren't part of any RFC. +const ( + // versionInformationParameterID is RFC 9368 (QUIC Version Negotiation). + versionInformationParameterID transportParameterID = 0x11 + // googleInitialRTTParameterID carries a cached RTT for a known server. Only + // sent when such a value exists, so we never send it. + googleInitialRTTParameterID transportParameterID = 0x3127 + // googleConnectionOptionsParameterID carries four-character option tags. + googleConnectionOptionsParameterID transportParameterID = 0x3128 +) + +// chromeConnectionOptions is the sole option tag sent in +// google_connection_options. +var chromeConnectionOptions = []byte{'O', 'R', 'I', 'G'} + +// chromeGREASEMaxValueLen bounds the GREASE transport parameter's value, whose +// length is drawn uniformly. Pinning it to one length would let a single +// connection exclude everything but this implementation. +const chromeGREASEMaxValueLen = 15 + +// marshalChrome encodes the client's transport parameters the way Chrome does. +// It differs from Marshal in four independently fingerprintable ways: +// +// 1. Parameters are written in a random order on every connection, rather than +// a fixed order with the GREASE value pinned to the front. +// 2. version_information and google_connection_options are sent, which quic-go +// does not send at all. +// 3. ack_delay_exponent, max_ack_delay, disable_active_migration and +// active_connection_id_limit are omitted, relying on the protocol defaults. +// 4. The GREASE parameter uses a large reserved id rather than a short one. +// +// Values come from p rather than being hardcoded, so what we advertise stays +// truthful about what this endpoint will honor; pinning the values themselves is +// Config.ChromeParrot's job. +// +// Omitting parameters is safe: the peer falls back to the same defaults it would +// assume for the client being imitated. +func (p *TransportParameters) marshalChrome() []byte { + // Collect each parameter as an independently encoded blob so we can shuffle + // them before concatenating. + params := make([][]byte, 0, 13) + add := func(b []byte) { params = append(params, b) } + + add(p.marshalVarintParam(nil, maxIdleTimeoutParameterID, uint64(p.MaxIdleTimeout/time.Millisecond))) + add(p.marshalVarintParam(nil, maxUDPPayloadSizeParameterID, uint64(p.MaxUDPPayloadSize))) + add(p.marshalVarintParam(nil, initialMaxDataParameterID, uint64(p.InitialMaxData))) + add(p.marshalVarintParam(nil, initialMaxStreamDataBidiLocalParameterID, uint64(p.InitialMaxStreamDataBidiLocal))) + add(p.marshalVarintParam(nil, initialMaxStreamDataBidiRemoteParameterID, uint64(p.InitialMaxStreamDataBidiRemote))) + add(p.marshalVarintParam(nil, initialMaxStreamDataUniParameterID, uint64(p.InitialMaxStreamDataUni))) + add(p.marshalVarintParam(nil, initialMaxStreamsBidiParameterID, uint64(p.MaxBidiStreamNum))) + add(p.marshalVarintParam(nil, initialMaxStreamsUniParameterID, uint64(p.MaxUniStreamNum))) + + // initial_source_connection_id. Normally empty, but encode whatever we have. + b := quicvarint.Append(nil, uint64(initialSourceConnectionIDParameterID)) + b = quicvarint.Append(b, uint64(p.InitialSourceConnectionID.Len())) + add(append(b, p.InitialSourceConnectionID.Bytes()...)) + + add(marshalChromeVersionInformation()) + + if p.MaxDatagramFrameSize != protocol.InvalidByteCount { + add(p.marshalVarintParam(nil, maxDatagramFrameSizeParameterID, uint64(p.MaxDatagramFrameSize))) + } + + b = quicvarint.Append(nil, uint64(googleConnectionOptionsParameterID)) + b = quicvarint.Append(b, uint64(len(chromeConnectionOptions))) + add(append(b, chromeConnectionOptions...)) + + add(marshalChromeGREASE()) + + shuffleParams(params) + + out := make([]byte, 0, 256) + for _, param := range params { + out = append(out, param...) + } + return out +} + +// marshalChromeVersionInformation encodes the version_information transport +// parameter (RFC 9368): chosen version 1, then an available-version list holding +// version 1 with one GREASE version inserted at a uniformly random index. A +// fixed order is both a statistical tell and, in one direction, a +// per-connection one. +func marshalChromeVersionInformation() []byte { + b := quicvarint.Append(nil, uint64(versionInformationParameterID)) + b = quicvarint.Append(b, 12) // 3 versions, 4 bytes each + b = binary.BigEndian.AppendUint32(b, uint32(protocol.Version1)) + + available := [2]uint32{uint32(protocol.Version1), greaseQUICVersion()} + if randUint64n(2) == 1 { + available[0], available[1] = available[1], available[0] + } + for _, v := range available { + b = binary.BigEndian.AppendUint32(b, v) + } + return b +} + +// greaseQUICVersion returns a random reserved QUIC version. RFC 9000 section 15 +// reserves versions matching 0x?a?a?a?a to exercise version negotiation. +func greaseQUICVersion() uint32 { + var buf [4]byte + rand.Read(buf[:]) + for i := range buf { + buf[i] = buf[i]&0xf0 | 0x0a + } + return binary.BigEndian.Uint32(buf[:]) +} + +// marshalChromeGREASE encodes a GREASE transport parameter with a large reserved +// id and a random value whose length varies per connection. +// +// RFC 9000 section 18.1 reserves ids of the form 31*N+27. N must be large enough +// that the id fills a full 8-byte varint, rather than the short id stock quic-go +// produces. +func marshalChromeGREASE() []byte { + // Keep the id in the range that encodes as a full 8-byte varint. + const nMin = (1<<30 - 27 + 30) / 31 // ceil((2^30-27)/31) + const nMax = (1<<62 - 1 - 27) / 31 // floor((2^62-1-27)/31) + n := nMin + randUint64n(nMax-nMin) + + b := quicvarint.Append(nil, 31*n+27) + valLen := randUint64n(chromeGREASEMaxValueLen + 1) + b = quicvarint.Append(b, valLen) + + val := make([]byte, valLen) + rand.Read(val) + return append(b, val...) +} + +// shuffleParams performs a Fisher-Yates shuffle backed by crypto/rand. +func shuffleParams(params [][]byte) { + for i := len(params) - 1; i > 0; i-- { + j := randUint64n(uint64(i + 1)) + params[i], params[j] = params[j], params[i] + } +} + +// randUint64n returns a random value in [0, n). The modulo bias is negligible +// at the magnitudes used here. +func randUint64n(n uint64) uint64 { + var buf [8]byte + rand.Read(buf[:]) + return binary.BigEndian.Uint64(buf[:]) % n +} diff --git a/third_party/quic-go/internal/wire/transport_parameters_chrome_test.go b/third_party/quic-go/internal/wire/transport_parameters_chrome_test.go new file mode 100644 index 0000000..2ddbd49 --- /dev/null +++ b/third_party/quic-go/internal/wire/transport_parameters_chrome_test.go @@ -0,0 +1,247 @@ +package wire + +import ( + "encoding/binary" + "strings" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/quicvarint" + "github.com/stretchr/testify/require" +) + +// chromeTestParams mirrors the values a Chrome-parroting client advertises. +func chromeTestParams() *TransportParameters { + return &TransportParameters{ + MaxIdleTimeout: 30 * time.Second, + MaxUDPPayloadSize: 1472, + InitialMaxData: 15728640, + InitialMaxStreamDataBidiLocal: 6291456, + InitialMaxStreamDataBidiRemote: 6291456, + InitialMaxStreamDataUni: 6291456, + MaxBidiStreamNum: 100, + MaxUniStreamNum: 103, + InitialSourceConnectionID: protocol.ConnectionID{}, + MaxDatagramFrameSize: 65536, + ChromeFingerprint: true, + } +} + +// parseParamIDs walks a marshalled transport parameter blob and returns the +// parameter ids in wire order, along with each one's value length. +func parseParamIDs(t *testing.T, b []byte) ([]uint64, map[uint64]uint64) { + t.Helper() + var ids []uint64 + lengths := make(map[uint64]uint64) + r := b + for len(r) > 0 { + id, n, err := quicvarint.Parse(r) + require.NoError(t, err) + r = r[n:] + l, n, err := quicvarint.Parse(r) + require.NoError(t, err) + r = r[n:] + require.GreaterOrEqual(t, uint64(len(r)), l, "truncated parameter value") + r = r[l:] + ids = append(ids, id) + lengths[id] = l + } + return ids, lengths +} + +func TestChromeTransportParametersSet(t *testing.T) { + ids, lengths := parseParamIDs(t, chromeTestParams().Marshal(protocol.PerspectiveClient)) + + // The full parameter set and nothing else. The GREASE parameter has a random + // id, so it is checked separately. + expected := []uint64{ + 0x01, // max_idle_timeout + 0x03, // max_udp_payload_size + 0x04, // initial_max_data + 0x05, // initial_max_stream_data_bidi_local + 0x06, // initial_max_stream_data_bidi_remote + 0x07, // initial_max_stream_data_uni + 0x08, // initial_max_streams_bidi + 0x09, // initial_max_streams_uni + 0x0f, // initial_source_connection_id + 0x11, // version_information + 0x20, // max_datagram_frame_size + 0x3128, + } + for _, id := range expected { + require.Contains(t, ids, id, "missing parameter 0x%x", id) + } + require.Len(t, ids, len(expected)+1, "unexpected parameter count (expected Chrome's set plus one GREASE)") + + // Parameters that must never be sent; emitting one gives the game away. + for _, id := range []uint64{ + 0x0a, // ack_delay_exponent + 0x0b, // max_ack_delay + 0x0c, // disable_active_migration + 0x0e, // active_connection_id_limit + 0x1d, // reset_stream_at + } { + require.NotContains(t, ids, id, "parameter 0x%x should be absent", id) + } + + // The source connection ID is zero-length. + require.Zero(t, lengths[0x0f]) +} + +func TestChromeTransportParametersGREASE(t *testing.T) { + ids, lengths := parseParamIDs(t, chromeTestParams().Marshal(protocol.PerspectiveClient)) + + var greaseID uint64 + for _, id := range ids { + if id > 0x3128 { + greaseID = id + break + } + } + require.NotZero(t, greaseID, "no GREASE parameter found") + + // RFC 9000 section 18.1 reserves ids of the form 31*N+27. + require.Equal(t, uint64(27), greaseID%31, "GREASE id must be 31*N+27") + // The id must need a full 8-byte varint, unlike quic-go's short one. + require.Equal(t, 8, quicvarint.Len(greaseID)) + require.LessOrEqual(t, lengths[greaseID], uint64(chromeGREASEMaxValueLen)) +} + +func TestChromeTransportParametersGREASEValueLengthVaries(t *testing.T) { + // The GREASE value length is drawn uniformly. A fixed length would be a + // per-connection exclusion whenever the imitated client picked another. + seen := make(map[uint64]struct{}) + for range 300 { + ids, lengths := parseParamIDs(t, chromeTestParams().Marshal(protocol.PerspectiveClient)) + for _, id := range ids { + if id > 0x3128 { + seen[lengths[id]] = struct{}{} + break + } + } + } + // Far fewer distinct lengths than draws would mean the range is too narrow. + require.Greater(t, len(seen), 12, "GREASE value length barely varies: %v", seen) + for l := range seen { + require.LessOrEqual(t, l, uint64(chromeGREASEMaxValueLen)) + } +} + +func TestChromeTransportParametersVersionOrderIsShuffled(t *testing.T) { + // The available-version list is shuffled; a fixed order is a tell. + var v1First, greaseFirst int + for range 200 { + b := chromeTestParams().Marshal(protocol.PerspectiveClient) + r := b + for len(r) > 0 { + id, n, err := quicvarint.Parse(r) + require.NoError(t, err) + r = r[n:] + l, n, err := quicvarint.Parse(r) + require.NoError(t, err) + r = r[n:] + if id == uint64(versionInformationParameterID) { + if binary.BigEndian.Uint32(r[4:8]) == uint32(protocol.Version1) { + v1First++ + } else { + greaseFirst++ + } + break + } + r = r[l:] + } + } + require.Positive(t, v1First, "version 1 never leads the available list") + require.Positive(t, greaseFirst, "the GREASE version never leads the available list") +} + +func TestChromeTransportParametersVersionInformation(t *testing.T) { + b := chromeTestParams().Marshal(protocol.PerspectiveClient) + + // Locate version_information and decode its value. + r := b + var val []byte + for len(r) > 0 { + id, n, err := quicvarint.Parse(r) + require.NoError(t, err) + r = r[n:] + l, n, err := quicvarint.Parse(r) + require.NoError(t, err) + r = r[n:] + if id == uint64(versionInformationParameterID) { + val = r[:l] + break + } + r = r[l:] + } + require.Len(t, val, 12, "chosen version plus two available versions") + + // Chosen version is always 1. + require.Equal(t, uint32(protocol.Version1), binary.BigEndian.Uint32(val[0:4])) + + // The available list holds version 1 and one GREASE version in either order, + // so identify them by value rather than position. + first := binary.BigEndian.Uint32(val[4:8]) + second := binary.BigEndian.Uint32(val[8:12]) + var grease uint32 + switch { + case first == uint32(protocol.Version1): + grease = second + case second == uint32(protocol.Version1): + grease = first + default: + t.Fatalf("available versions %#x, %#x contain no version 1", first, second) + } + + // RFC 9000 section 15 reserves versions matching 0x?a?a?a?a. + for i := range 4 { + v := byte(grease >> (8 * i)) + require.Equal(t, byte(0x0a), v&0x0f, "GREASE version byte %d low nibble", i) + } +} + +func TestChromeTransportParametersOrderIsShuffled(t *testing.T) { + // The parameter order is permuted per connection; a fixed order is itself a + // fingerprint. Repeated runs colliding on one order would be vanishingly + // unlikely. + seen := make(map[string]struct{}) + for range 30 { + ids, _ := parseParamIDs(t, chromeTestParams().Marshal(protocol.PerspectiveClient)) + var key strings.Builder + for _, id := range ids { + // Normalize the random GREASE id so only its position matters. + if id > 0x3128 { + id = 0xffff + } + key.WriteString(string(rune(id)) + ",") + } + seen[key.String()] = struct{}{} + } + require.Greater(t, len(seen), 1, "parameter order is not being shuffled") +} + +func TestChromeTransportParametersRoundTrip(t *testing.T) { + // Whatever cosmetics we apply, the peer must still be able to parse our + // parameters and read back the values we actually intend to honor. + p := chromeTestParams() + var parsed TransportParameters + require.NoError(t, parsed.Unmarshal(p.Marshal(protocol.PerspectiveClient), protocol.PerspectiveClient)) + + require.Equal(t, p.MaxIdleTimeout, parsed.MaxIdleTimeout) + require.Equal(t, p.MaxUDPPayloadSize, parsed.MaxUDPPayloadSize) + require.Equal(t, p.InitialMaxData, parsed.InitialMaxData) + require.Equal(t, p.InitialMaxStreamDataBidiLocal, parsed.InitialMaxStreamDataBidiLocal) + require.Equal(t, p.InitialMaxStreamDataBidiRemote, parsed.InitialMaxStreamDataBidiRemote) + require.Equal(t, p.InitialMaxStreamDataUni, parsed.InitialMaxStreamDataUni) + require.Equal(t, p.MaxBidiStreamNum, parsed.MaxBidiStreamNum) + require.Equal(t, p.MaxUniStreamNum, parsed.MaxUniStreamNum) + require.Equal(t, p.MaxDatagramFrameSize, parsed.MaxDatagramFrameSize) + require.Equal(t, p.InitialSourceConnectionID, parsed.InitialSourceConnectionID) + + // The omitted parameters must come back as the protocol defaults. + require.Equal(t, protocol.DefaultAckDelayExponent, int(parsed.AckDelayExponent)) + require.Equal(t, protocol.DefaultMaxAckDelay, parsed.MaxAckDelay) + require.Equal(t, uint64(protocol.DefaultActiveConnectionIDLimit), parsed.ActiveConnectionIDLimit) + require.False(t, parsed.DisableActiveMigration) +} diff --git a/third_party/quic-go/internal/wire/version_negotiation.go b/third_party/quic-go/internal/wire/version_negotiation.go new file mode 100644 index 0000000..04a3ed8 --- /dev/null +++ b/third_party/quic-go/internal/wire/version_negotiation.go @@ -0,0 +1,53 @@ +package wire + +import ( + "crypto/rand" + "encoding/binary" + "errors" + + "github.com/apernet/quic-go/internal/protocol" +) + +// ParseVersionNegotiationPacket parses a Version Negotiation packet. +func ParseVersionNegotiationPacket(b []byte) (dest, src protocol.ArbitraryLenConnectionID, _ []protocol.Version, _ error) { + n, dest, src, err := ParseArbitraryLenConnectionIDs(b) + if err != nil { + return nil, nil, nil, err + } + b = b[n:] + if len(b) == 0 { + //nolint:staticcheck // SA1021: the packet is called Version Negotiation packet + return nil, nil, nil, errors.New("Version Negotiation packet has empty version list") + } + if len(b)%4 != 0 { + //nolint:staticcheck // SA1021: the packet is called Version Negotiation packet + return nil, nil, nil, errors.New("Version Negotiation packet has a version list with an invalid length") + } + versions := make([]protocol.Version, len(b)/4) + for i := 0; len(b) > 0; i++ { + versions[i] = protocol.Version(binary.BigEndian.Uint32(b[:4])) + b = b[4:] + } + return dest, src, versions, nil +} + +// ComposeVersionNegotiation composes a Version Negotiation +func ComposeVersionNegotiation(destConnID, srcConnID protocol.ArbitraryLenConnectionID, versions []protocol.Version) []byte { + greasedVersions := protocol.GetGreasedVersions(versions) + expectedLen := 1 /* type byte */ + 4 /* version field */ + 1 /* dest connection ID length field */ + destConnID.Len() + 1 /* src connection ID length field */ + srcConnID.Len() + len(greasedVersions)*4 + buf := make([]byte, 1+4 /* type byte and version field */, expectedLen) + _, _ = rand.Read(buf[:1]) // ignore the error here. It is not critical to have perfect random here. + // Setting the "QUIC bit" (0x40) is not required by the RFC, + // but it allows clients to demultiplex QUIC with a long list of other protocols. + // See RFC 9443 and https://mailarchive.ietf.org/arch/msg/quic/oR4kxGKY6mjtPC1CZegY1ED4beg/ for details. + buf[0] |= 0xc0 + // The next 4 bytes are left at 0 (version number). + buf = append(buf, uint8(destConnID.Len())) + buf = append(buf, destConnID.Bytes()...) + buf = append(buf, uint8(srcConnID.Len())) + buf = append(buf, srcConnID.Bytes()...) + for _, v := range greasedVersions { + buf = binary.BigEndian.AppendUint32(buf, uint32(v)) + } + return buf +} diff --git a/third_party/quic-go/internal/wire/version_negotiation_test.go b/third_party/quic-go/internal/wire/version_negotiation_test.go new file mode 100644 index 0000000..cc3bcd7 --- /dev/null +++ b/third_party/quic-go/internal/wire/version_negotiation_test.go @@ -0,0 +1,103 @@ +package wire + +import ( + "crypto/rand" + "encoding/binary" + mrand "math/rand/v2" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestParseVersionNegotiationPacket(t *testing.T) { + randConnID := func(l int) protocol.ArbitraryLenConnectionID { + b := make(protocol.ArbitraryLenConnectionID, l) + _, err := rand.Read(b) + require.NoError(t, err) + return b + } + + srcConnID := randConnID(mrand.IntN(255) + 1) + destConnID := randConnID(mrand.IntN(255) + 1) + versions := []protocol.Version{0x22334455, 0x33445566} + data := []byte{0x80, 0, 0, 0, 0} + data = append(data, uint8(len(destConnID))) + data = append(data, destConnID...) + data = append(data, uint8(len(srcConnID))) + data = append(data, srcConnID...) + for _, v := range versions { + data = append(data, []byte{0, 0, 0, 0}...) + binary.BigEndian.PutUint32(data[len(data)-4:], uint32(v)) + } + require.True(t, IsVersionNegotiationPacket(data)) + dest, src, supportedVersions, err := ParseVersionNegotiationPacket(data) + require.NoError(t, err) + require.Equal(t, destConnID, dest) + require.Equal(t, srcConnID, src) + require.Equal(t, versions, supportedVersions) +} + +func TestParseVersionNegotiationPacketWithInvalidLength(t *testing.T) { + connID := protocol.ArbitraryLenConnectionID{1, 2, 3, 4, 5, 6, 7, 8} + versions := []protocol.Version{0x22334455, 0x33445566} + data := ComposeVersionNegotiation(connID, connID, versions) + _, _, _, err := ParseVersionNegotiationPacket(data[:len(data)-2]) + require.EqualError(t, err, "Version Negotiation packet has a version list with an invalid length") +} + +func TestParseVersionNegotiationPacketEmptyVersions(t *testing.T) { + connID := protocol.ArbitraryLenConnectionID{1, 2, 3, 4, 5, 6, 7, 8} + versions := []protocol.Version{0x22334455} + data := ComposeVersionNegotiation(connID, connID, versions) + // remove 8 bytes (two versions), since ComposeVersionNegotiation also added a reserved version number + data = data[:len(data)-8] + _, _, _, err := ParseVersionNegotiationPacket(data) + require.EqualError(t, err, "Version Negotiation packet has empty version list") +} + +func TestComposeVersionNegotiationWithReservedVersion(t *testing.T) { + srcConnID := protocol.ArbitraryLenConnectionID{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0x13, 0x37} + destConnID := protocol.ArbitraryLenConnectionID{1, 2, 3, 4, 5, 6, 7, 8} + versions := []protocol.Version{1001, 1003} + data := ComposeVersionNegotiation(destConnID, srcConnID, versions) + require.True(t, IsLongHeaderPacket(data[0])) + require.NotZero(t, data[0]&0x40) + v, err := ParseVersion(data) + require.NoError(t, err) + require.Zero(t, v) + dest, src, supportedVersions, err := ParseVersionNegotiationPacket(data) + require.NoError(t, err) + require.Equal(t, destConnID, dest) + require.Equal(t, srcConnID, src) + // the supported versions should include one reserved version number + require.Len(t, supportedVersions, len(versions)+1) + for _, v := range versions { + require.Contains(t, supportedVersions, v) + } + var reservedVersion protocol.Version +versionLoop: + for _, ver := range supportedVersions { + for _, v := range versions { + if v == ver { + continue versionLoop + } + } + reservedVersion = ver + } + require.NotZero(t, reservedVersion) + require.True(t, reservedVersion&0x0f0f0f0f == 0x0a0a0a0a) // check that it's a greased version number +} + +func BenchmarkComposeVersionNegotiationPacket(b *testing.B) { + b.ReportAllocs() + + supportedVersions := []protocol.Version{protocol.Version2, protocol.Version1, 0x1337} + destConnID := protocol.ArbitraryLenConnectionID{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 0xa, 0xb, 0xc, 0xd} + srcConnID := protocol.ArbitraryLenConnectionID{10, 9, 8, 7, 6, 5, 4, 3, 2, 1} + + for b.Loop() { + ComposeVersionNegotiation(destConnID, srcConnID, supportedVersions) + } +} diff --git a/third_party/quic-go/interop/Dockerfile b/third_party/quic-go/interop/Dockerfile new file mode 100644 index 0000000..ccc1c92 --- /dev/null +++ b/third_party/quic-go/interop/Dockerfile @@ -0,0 +1,38 @@ +FROM martenseemann/quic-network-simulator-endpoint:latest AS builder + +ARG TARGETPLATFORM +RUN echo "TARGETPLATFORM: ${TARGETPLATFORM}" + +RUN apt-get update && apt-get install -y wget tar git && rm -rf /var/lib/apt/lists/* + +ENV GOVERSION=1.26.0 + +RUN platform=$(echo ${TARGETPLATFORM} | tr '/' '-') && \ + filename="go${GOVERSION}.${platform}.tar.gz" && \ + wget https://dl.google.com/go/${filename} && \ + tar xfz ${filename} && \ + rm ${filename} + +ENV PATH="/go/bin:${PATH}" + +# build with --build-arg CACHEBUST=$(date +%s) +ARG CACHEBUST=1 + +COPY . /quic-go +WORKDIR /quic-go + +RUN git rev-parse HEAD | tee commit.txt +RUN go build -o server -ldflags="-X github.com/quic-go/quic-go/qlog.quicGoVersion=$(git describe --always --long --dirty)" interop/server/main.go +RUN go build -o client -ldflags="-X github.com/quic-go/quic-go/qlog.quicGoVersion=$(git describe --always --long --dirty)" interop/client/main.go + + +FROM martenseemann/quic-network-simulator-endpoint:latest + +WORKDIR /quic-go + +COPY --from=builder /quic-go/commit.txt /quic-go/server /quic-go/client ./ +COPY --from=builder /quic-go/interop/run_endpoint.sh ./ + +RUN chmod +x run_endpoint.sh + +ENTRYPOINT [ "./run_endpoint.sh" ] diff --git a/third_party/quic-go/interop/client/main.go b/third_party/quic-go/interop/client/main.go new file mode 100644 index 0000000..9d72aab --- /dev/null +++ b/third_party/quic-go/interop/client/main.go @@ -0,0 +1,210 @@ +package main + +import ( + "crypto/tls" + "errors" + "flag" + "fmt" + "io" + "log" + "net/http" + "os" + "strings" + "time" + + "golang.org/x/sync/errgroup" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qtls" + "github.com/apernet/quic-go/interop/http09" + "github.com/apernet/quic-go/interop/utils" +) + +var errUnsupported = errors.New("unsupported test case") + +var tlsConf *tls.Config + +func main() { + logFile, err := os.Create("/logs/log.txt") + if err != nil { + fmt.Printf("Could not create log file: %s\n", err.Error()) + os.Exit(1) + } + defer logFile.Close() + log.SetOutput(logFile) + + keyLog, err := utils.GetSSLKeyLog() + if err != nil { + fmt.Printf("Could not create key log: %s\n", err.Error()) + os.Exit(1) + } + if keyLog != nil { + defer keyLog.Close() + } + + tlsConf = &tls.Config{ + InsecureSkipVerify: true, + KeyLogWriter: keyLog, + } + testcase := os.Getenv("TESTCASE") + if err := runTestcase(testcase); err != nil { + if err == errUnsupported { + fmt.Printf("unsupported test case: %s\n", testcase) + os.Exit(127) + } + fmt.Printf("Downloading files failed: %s\n", err.Error()) + os.Exit(1) + } +} + +func runTestcase(testcase string) error { + flag.Parse() + urls := flag.Args() + + quicConf := &quic.Config{Tracer: utils.NewQLOGConnectionTracer} + + if testcase == "http3" { + r := &http3.Transport{ + TLSClientConfig: tlsConf, + QUICConfig: quicConf, + } + defer r.Close() + return downloadFiles(r, urls, false) + } + + r := &http09.RoundTripper{ + TLSClientConfig: tlsConf, + QuicConfig: quicConf, + } + defer r.Close() + + switch testcase { + case "handshake", "transfer", "retry": + case "keyupdate": + handshake.FirstKeyUpdateInterval = 100 + case "chacha20": + reset := qtls.SetCipherSuite(tls.TLS_CHACHA20_POLY1305_SHA256) + defer reset() + case "multiconnect": + return runMultiConnectTest(r, urls) + case "versionnegotiation": + return runVersionNegotiationTest(r, urls) + case "resumption": + return runResumptionTest(r, urls, false) + case "zerortt": + return runResumptionTest(r, urls, true) + default: + return errUnsupported + } + + return downloadFiles(r, urls, false) +} + +func runVersionNegotiationTest(r *http09.RoundTripper, urls []string) error { + if len(urls) != 1 { + return errors.New("expected at least 2 URLs") + } + protocol.SupportedVersions = []protocol.Version{0x1a2a3a4a} + err := downloadFile(r, urls[0], false) + if err == nil { + return errors.New("expected version negotiation to fail") + } + if !strings.Contains(err.Error(), "No compatible QUIC version found") { + return fmt.Errorf("expect version negotiation error, got: %s", err.Error()) + } + return nil +} + +func runMultiConnectTest(r *http09.RoundTripper, urls []string) error { + for _, url := range urls { + if err := downloadFile(r, url, false); err != nil { + return err + } + if err := r.Close(); err != nil { + return err + } + } + return nil +} + +type sessionCache struct { + tls.ClientSessionCache + put chan<- struct{} +} + +func newSessionCache(c tls.ClientSessionCache) (tls.ClientSessionCache, <-chan struct{}) { + put := make(chan struct{}, 100) + return &sessionCache{ClientSessionCache: c, put: put}, put +} + +func (c *sessionCache) Put(key string, cs *tls.ClientSessionState) { + c.ClientSessionCache.Put(key, cs) + c.put <- struct{}{} +} + +func runResumptionTest(r *http09.RoundTripper, urls []string, use0RTT bool) error { + if len(urls) < 2 { + return errors.New("expected at least 2 URLs") + } + + var put <-chan struct{} + tlsConf.ClientSessionCache, put = newSessionCache(tls.NewLRUClientSessionCache(1)) + + // do the first transfer + if err := downloadFiles(r, urls[:1], false); err != nil { + return err + } + + // wait for the session ticket to arrive + select { + case <-time.NewTimer(10 * time.Second).C: + return errors.New("expected to receive a session ticket within 10 seconds") + case <-put: + } + + if err := r.Close(); err != nil { + return err + } + + // reestablish the connection, using the session ticket that the server (hopefully provided) + defer r.Close() + return downloadFiles(r, urls[1:], use0RTT) +} + +func downloadFiles(cl http.RoundTripper, urls []string, use0RTT bool) error { + var g errgroup.Group + for _, u := range urls { + url := u + g.Go(func() error { + return downloadFile(cl, url, use0RTT) + }) + } + return g.Wait() +} + +func downloadFile(cl http.RoundTripper, url string, use0RTT bool) error { + method := http.MethodGet + if use0RTT { + method = http09.MethodGet0RTT + } + req, err := http.NewRequest(method, url, nil) + if err != nil { + return err + } + rsp, err := cl.RoundTrip(req) + if err != nil { + return err + } + defer rsp.Body.Close() + + file, err := os.Create("/downloads" + req.URL.Path) + if err != nil { + return err + } + defer file.Close() + _, err = io.Copy(file, rsp.Body) + return err +} diff --git a/third_party/quic-go/interop/http09/client.go b/third_party/quic-go/interop/http09/client.go new file mode 100644 index 0000000..103e4c9 --- /dev/null +++ b/third_party/quic-go/interop/http09/client.go @@ -0,0 +1,159 @@ +package http09 + +import ( + "context" + "crypto/tls" + "errors" + "io" + "log" + "net" + "net/http" + "strings" + "sync" + + "golang.org/x/net/idna" + + "github.com/apernet/quic-go" +) + +// MethodGet0RTT allows a GET request to be sent using 0-RTT. +// Note that 0-RTT data doesn't provide replay protection. +const MethodGet0RTT = "GET_0RTT" + +// RoundTripper performs HTTP/0.9 roundtrips over QUIC. +type RoundTripper struct { + mutex sync.Mutex + + TLSClientConfig *tls.Config + QuicConfig *quic.Config + + clients map[string]*client +} + +var _ http.RoundTripper = &RoundTripper{} + +// RoundTrip performs a HTTP/0.9 request. +// It only supports GET requests. +func (r *RoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + if req.Method != http.MethodGet && req.Method != MethodGet0RTT { + return nil, errors.New("only GET requests supported") + } + + log.Printf("Requesting %s.\n", req.URL) + + r.mutex.Lock() + hostname := authorityAddr("https", hostnameFromRequest(req)) + if r.clients == nil { + r.clients = make(map[string]*client) + } + c, ok := r.clients[hostname] + if !ok { + tlsConf := &tls.Config{} + if r.TLSClientConfig != nil { + tlsConf = r.TLSClientConfig.Clone() + } + tlsConf.NextProtos = []string{NextProto} + c = &client{ + hostname: hostname, + tlsConf: tlsConf, + quicConf: r.QuicConfig, + } + r.clients[hostname] = c + } + r.mutex.Unlock() + return c.RoundTrip(req) +} + +// Close closes the roundtripper. +func (r *RoundTripper) Close() error { + r.mutex.Lock() + defer r.mutex.Unlock() + + for id, c := range r.clients { + if err := c.Close(); err != nil { + return err + } + delete(r.clients, id) + } + return nil +} + +type client struct { + hostname string + tlsConf *tls.Config + quicConf *quic.Config + + once sync.Once + conn *quic.Conn + dialErr error +} + +func (c *client) RoundTrip(req *http.Request) (*http.Response, error) { + c.once.Do(func() { + c.conn, c.dialErr = quic.DialAddrEarly(context.Background(), c.hostname, c.tlsConf, c.quicConf) + }) + if c.dialErr != nil { + return nil, c.dialErr + } + if req.Method != MethodGet0RTT { + <-c.conn.HandshakeComplete() + } + return c.doRequest(req) +} + +func (c *client) doRequest(req *http.Request) (*http.Response, error) { + str, err := c.conn.OpenStreamSync(context.Background()) + if err != nil { + return nil, err + } + cmd := "GET " + req.URL.Path + "\r\n" + if _, err := str.Write([]byte(cmd)); err != nil { + return nil, err + } + if err := str.Close(); err != nil { + return nil, err + } + rsp := &http.Response{ + Proto: "HTTP/0.9", + ProtoMajor: 0, + ProtoMinor: 9, + Request: req, + Body: io.NopCloser(str), + } + return rsp, nil +} + +func (c *client) Close() error { + if c.conn == nil { + return nil + } + return c.conn.CloseWithError(0, "") +} + +func hostnameFromRequest(req *http.Request) string { + if req.URL != nil { + return req.URL.Host + } + return "" +} + +// authorityAddr returns a given authority (a host/IP, or host:port / ip:port) +// and returns a host:port. The port 443 is added if needed. +func authorityAddr(scheme string, authority string) (addr string) { + host, port, err := net.SplitHostPort(authority) + if err != nil { // authority didn't have a port + port = "443" + if scheme == "http" { + port = "80" + } + host = authority + } + if a, err := idna.ToASCII(host); err == nil { + host = a + } + // IPv6 address literal, without a port: + if strings.HasPrefix(host, "[") && strings.HasSuffix(host, "]") { + return host + ":" + port + } + return net.JoinHostPort(host, port) +} diff --git a/third_party/quic-go/interop/http09/http_test.go b/third_party/quic-go/interop/http09/http_test.go new file mode 100644 index 0000000..ad83b3e --- /dev/null +++ b/third_party/quic-go/interop/http09/http_test.go @@ -0,0 +1,77 @@ +package http09 + +import ( + "crypto/tls" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "testing" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/internal/testdata" + + "github.com/stretchr/testify/require" +) + +func startServer(t *testing.T) net.Addr { + t.Helper() + server := &Server{} + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + tr := &quic.Transport{Conn: conn} + tlsConf := testdata.GetTLSConfig() + tlsConf.NextProtos = []string{NextProto} + ln, err := tr.ListenEarly(tlsConf, &quic.Config{}) + require.NoError(t, err) + done := make(chan struct{}) + go func() { + defer close(done) + _ = server.ServeListener(ln) + }() + t.Cleanup(func() { + require.NoError(t, ln.Close()) + <-done + }) + return ln.Addr() +} + +func TestHTTPRequest(t *testing.T) { + http.HandleFunc("/helloworld", func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte("Hello World!")) + }) + + addr := startServer(t) + + rt := &RoundTripper{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}} + t.Cleanup(func() { rt.Close() }) + + req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("https://%s/helloworld", addr), nil) + rsp, err := rt.RoundTrip(req) + require.NoError(t, err) + data, err := io.ReadAll(rsp.Body) + require.NoError(t, err) + require.Equal(t, []byte("Hello World!"), data) +} + +func TestHTTPHeaders(t *testing.T) { + http.HandleFunc("/headers", func(w http.ResponseWriter, r *http.Request) { + w.Header().Add("foo", "bar") + w.WriteHeader(1337) + _, _ = w.Write([]byte("done")) + }) + + addr := startServer(t) + + rt := &RoundTripper{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}} + t.Cleanup(func() { rt.Close() }) + + req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("https://%s/headers", addr), nil) + rsp, err := rt.RoundTrip(req) + require.NoError(t, err) + data, err := io.ReadAll(rsp.Body) + require.NoError(t, err) + require.Equal(t, []byte("done"), data) + // HTTP/0.9 doesn't support HTTP headers +} diff --git a/third_party/quic-go/interop/http09/server.go b/third_party/quic-go/interop/http09/server.go new file mode 100644 index 0000000..b6f4004 --- /dev/null +++ b/third_party/quic-go/interop/http09/server.go @@ -0,0 +1,121 @@ +package http09 + +import ( + "context" + "io" + "log" + "net/http" + "net/url" + "runtime" + "strings" + + "github.com/apernet/quic-go" +) + +const NextProto = "hq-interop" + +type responseWriter struct { + io.Writer + headers http.Header +} + +var _ http.ResponseWriter = &responseWriter{} + +func (w *responseWriter) Header() http.Header { + if w.headers == nil { + w.headers = make(http.Header) + } + return w.headers +} + +func (w *responseWriter) WriteHeader(int) {} + +// Server is a HTTP/0.9 server listening for QUIC connections. +type Server struct { + Handler *http.ServeMux +} + +// ServeListener serves HTTP/0.9 on all connections accepted from a QUIC listener. +func (s *Server) ServeListener(ln *quic.EarlyListener) error { + for { + conn, err := ln.Accept(context.Background()) + if err != nil { + return err + } + go s.handleConn(conn) + } +} + +func (s *Server) handleConn(conn *quic.Conn) { + for { + str, err := conn.AcceptStream(context.Background()) + if err != nil { + log.Printf("Error accepting stream: %s\n", err.Error()) + return + } + go func() { + if err := s.handleStream(str); err != nil { + log.Printf("Handling stream failed: %s\n", err.Error()) + } + }() + } +} + +func (s *Server) handleStream(str *quic.Stream) error { + reqBytes, err := io.ReadAll(str) + if err != nil { + return err + } + request := string(reqBytes) + request = strings.TrimRight(request, "\r\n") + request = strings.TrimRight(request, " ") + + log.Printf("Received request: %s\n", request) + + if request[:5] != "GET /" { + str.CancelWrite(42) + return nil + } + + u, err := url.Parse(request[4:]) + if err != nil { + return err + } + u.Scheme = "https" + + req := &http.Request{ + Method: http.MethodGet, + Proto: "HTTP/0.9", + ProtoMajor: 0, + ProtoMinor: 9, + Body: str, + URL: u, + } + + handler := s.Handler + if handler == nil { + handler = http.DefaultServeMux + } + + var panicked bool + func() { + defer func() { + if p := recover(); p != nil { + // Copied from net/http/server.go + const size = 64 << 10 + buf := make([]byte, size) + buf = buf[:runtime.Stack(buf, false)] + log.Printf("http: panic serving: %v\n%s", p, buf) + panicked = true + } + }() + handler.ServeHTTP(&responseWriter{Writer: str}, req) + }() + + if panicked { + if _, err := str.Write([]byte("500")); err != nil { + return err + } + } + return str.Close() +} diff --git a/third_party/quic-go/interop/run_endpoint.sh b/third_party/quic-go/interop/run_endpoint.sh new file mode 100644 index 0000000..9c3ee55 --- /dev/null +++ b/third_party/quic-go/interop/run_endpoint.sh @@ -0,0 +1,19 @@ +#!/bin/bash +set -e + +# Set up the routing needed for the simulation. +/setup.sh + +echo "Using commit:" `cat commit.txt` + +if [ "$ROLE" == "client" ]; then + # Wait for the simulator to start up. + /wait-for-it.sh sim:57832 -s -t 10 + echo "Starting QUIC client..." + echo "Client params: $CLIENT_PARAMS" + echo "Test case: $TESTCASE" + QUIC_GO_LOG_LEVEL=debug ./client $CLIENT_PARAMS $REQUESTS +else + echo "Running QUIC server." + QUIC_GO_LOG_LEVEL=debug ./server "$@" +fi diff --git a/third_party/quic-go/interop/server/main.go b/third_party/quic-go/interop/server/main.go new file mode 100644 index 0000000..f5c6a71 --- /dev/null +++ b/third_party/quic-go/interop/server/main.go @@ -0,0 +1,105 @@ +package main + +import ( + "crypto/tls" + "fmt" + "log" + "net" + "net/http" + "os" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/internal/qtls" + "github.com/apernet/quic-go/interop/http09" + "github.com/apernet/quic-go/interop/utils" +) + +func main() { + logFile, err := os.Create("/logs/log.txt") + if err != nil { + fmt.Printf("Could not create log file: %s\n", err.Error()) + os.Exit(1) + } + defer logFile.Close() + log.SetOutput(logFile) + + keyLog, err := utils.GetSSLKeyLog() + if err != nil { + fmt.Printf("Could not create key log: %s\n", err.Error()) + os.Exit(1) + } + if keyLog != nil { + defer keyLog.Close() + } + + testcase := os.Getenv("TESTCASE") + + quicConf := &quic.Config{ + Allow0RTT: testcase == "zerortt", + Tracer: utils.NewQLOGConnectionTracer, + } + cert, err := tls.LoadX509KeyPair("/certs/cert.pem", "/certs/priv.key") + if err != nil { + fmt.Println(err) + os.Exit(1) + } + tlsConf := &tls.Config{ + Certificates: []tls.Certificate{cert}, + KeyLogWriter: keyLog, + NextProtos: []string{http09.NextProto}, + } + + switch testcase { + case "versionnegotiation", "handshake", "retry", "transfer", "resumption", "multiconnect", "zerortt": + err = runHTTP09Server(tlsConf, quicConf, testcase == "retry") + case "chacha20": + reset := qtls.SetCipherSuite(tls.TLS_CHACHA20_POLY1305_SHA256) + defer reset() + err = runHTTP09Server(tlsConf, quicConf, false) + case "http3": + tlsConf.NextProtos = []string{http3.NextProtoH3} + err = runHTTP3Server(tlsConf, quicConf) + default: + fmt.Printf("unsupported test case: %s\n", testcase) + os.Exit(127) + } + + if err != nil { + fmt.Printf("Error running server: %s\n", err.Error()) + os.Exit(1) + } +} + +func runHTTP09Server(tlsConf *tls.Config, quicConf *quic.Config, forceRetry bool) error { + http.DefaultServeMux.Handle("/", http.FileServer(http.Dir("/www"))) + server := http09.Server{} + + udpAddr, err := net.ResolveUDPAddr("udp", ":443") + if err != nil { + return err + } + conn, err := net.ListenUDP("udp", udpAddr) + if err != nil { + return err + } + tr := &quic.Transport{ + Conn: conn, + VerifySourceAddress: func(net.Addr) bool { return forceRetry }, + } + ln, err := tr.ListenEarly(tlsConf, quicConf) + if err != nil { + return err + } + return server.ServeListener(ln) +} + +func runHTTP3Server(tlsConf *tls.Config, quicConf *quic.Config) error { + server := http3.Server{ + Addr: ":443", + TLSConfig: tlsConf, + QUICConfig: quicConf, + } + http.DefaultServeMux.Handle("/", http.FileServer(http.Dir("/www"))) + return server.ListenAndServe() +} diff --git a/third_party/quic-go/interop/utils/logging.go b/third_party/quic-go/interop/utils/logging.go new file mode 100644 index 0000000..2e39012 --- /dev/null +++ b/third_party/quic-go/interop/utils/logging.go @@ -0,0 +1,58 @@ +package utils + +import ( + "bufio" + "context" + "fmt" + "io" + "log" + "os" + "strings" + + "github.com/apernet/quic-go" + h3qlog "github.com/apernet/quic-go/http3/qlog" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +// GetSSLKeyLog creates a file for the TLS key log +func GetSSLKeyLog() (io.WriteCloser, error) { + filename := os.Getenv("SSLKEYLOGFILE") + if len(filename) == 0 { + return nil, nil + } + f, err := os.Create(filename) + if err != nil { + return nil, err + } + return f, nil +} + +// NewQLOGConnectionTracer create a qlog file in QLOGDIR +func NewQLOGConnectionTracer(_ context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace { + qlogDir := os.Getenv("QLOGDIR") + if len(qlogDir) == 0 { + return nil + } + if _, err := os.Stat(qlogDir); os.IsNotExist(err) { + if err := os.MkdirAll(qlogDir, 0o666); err != nil { + log.Fatalf("failed to create qlog dir %s: %v", qlogDir, err) + } + } + path := fmt.Sprintf("%s/%s.sqlog", strings.TrimRight(qlogDir, "/"), connID) + f, err := os.Create(path) + if err != nil { + log.Printf("Failed to create qlog file %s: %s", path, err.Error()) + return nil + } + log.Printf("Created qlog file: %s\n", path) + fileSeq := qlogwriter.NewConnectionFileSeq( + utils.NewBufferedWriteCloser(bufio.NewWriter(f), f), + isClient, + connID, + []string{qlog.EventSchema, h3qlog.EventSchema}, + ) + go fileSeq.Run() + return fileSeq +} diff --git a/third_party/quic-go/metrics/dashboards/README.md b/third_party/quic-go/metrics/dashboards/README.md new file mode 100644 index 0000000..ace507b --- /dev/null +++ b/third_party/quic-go/metrics/dashboards/README.md @@ -0,0 +1,24 @@ +# quic-go Prometheus / Grafana Local Development Setup + +For local development and debugging, it can be useful to spin up a local Prometheus and Grafana instance. + +Please refer to the [documentation](https://quic-go.net/docs/quic/metrics/) for how to configure quic-go to expose Prometheus metrics. + +The configuration files in this directory assume that the application exposes the Prometheus endpoint at `http://localhost:5001/prometheus`: +```go +import "github.com/prometheus/client_golang/prometheus/promhttp" + +go func() { + http.Handle("/prometheus", promhttp.Handler()) + log.Fatal(http.ListenAndServe("localhost:5001", nil)) +}() +``` + +Prometheus and Grafana can be started using Docker Compose: + +Running: +```shell +docker compose up +``` + +[quic-go.json](./quic-go.json) contains the JSON model of an example Grafana dashboard. diff --git a/third_party/quic-go/metrics/dashboards/datasources.yml b/third_party/quic-go/metrics/dashboards/datasources.yml new file mode 100644 index 0000000..ed47ec1 --- /dev/null +++ b/third_party/quic-go/metrics/dashboards/datasources.yml @@ -0,0 +1,13 @@ +apiVersion: 1 + +deleteDatasources: + - name: Prometheus + orgId: 1 + +datasources: + - name: Prometheus + orgId: 1 + type: prometheus + access: proxy + url: http://prometheus:9090 + editable: false diff --git a/third_party/quic-go/metrics/dashboards/docker-compose.yml b/third_party/quic-go/metrics/dashboards/docker-compose.yml new file mode 100644 index 0000000..8e3a8a5 --- /dev/null +++ b/third_party/quic-go/metrics/dashboards/docker-compose.yml @@ -0,0 +1,25 @@ +version: '3.8' + +volumes: + prometheus_data: {} + grafana_data: {} + +services: + prometheus: + image: prom/prometheus:latest + container_name: prometheus + volumes: + - ./prometheus.yml:/etc/prometheus/prometheus.yml + - prometheus_data:/prometheus + command: + - '--config.file=/etc/prometheus/prometheus.yml' + expose: + - 9090 + grafana: + image: grafana/grafana:latest + container_name: grafana + volumes: + - grafana_data:/var/lib/grafana + - ./datasources.yml:/etc/grafana/provisioning/datasources/prom.yml + ports: + - "3000:3000" diff --git a/third_party/quic-go/metrics/dashboards/prometheus.yml b/third_party/quic-go/metrics/dashboards/prometheus.yml new file mode 100644 index 0000000..5018097 --- /dev/null +++ b/third_party/quic-go/metrics/dashboards/prometheus.yml @@ -0,0 +1,9 @@ +global: + scrape_interval: 15s + +scrape_configs: + - job_name: 'quic-go' + scrape_interval: 15s + static_configs: + - targets: ['host.docker.internal:5001'] + metrics_path: '/prometheus' diff --git a/third_party/quic-go/metrics/dashboards/quic-go.json b/third_party/quic-go/metrics/dashboards/quic-go.json new file mode 100644 index 0000000..48c7fa5 --- /dev/null +++ b/third_party/quic-go/metrics/dashboards/quic-go.json @@ -0,0 +1,926 @@ +{ + "__inputs": [ + { + "name": "DS_PROMETHEUS", + "label": "Prometheus", + "description": "", + "type": "datasource", + "pluginId": "prometheus", + "pluginName": "Prometheus" + } + ], + "__elements": {}, + "__requires": [ + { + "type": "grafana", + "id": "grafana", + "name": "Grafana", + "version": "10.2.3" + }, + { + "type": "datasource", + "id": "prometheus", + "name": "Prometheus", + "version": "1.0.0" + }, + { + "type": "panel", + "id": "stat", + "name": "Stat", + "version": "" + }, + { + "type": "panel", + "id": "timeseries", + "name": "Time series", + "version": "" + } + ], + "annotations": { + "list": [ + { + "builtIn": 1, + "datasource": { + "type": "grafana", + "uid": "-- Grafana --" + }, + "enable": true, + "hide": true, + "iconColor": "rgba(0, 211, 255, 1)", + "name": "Annotations & Alerts", + "type": "dashboard" + } + ] + }, + "editable": true, + "fiscalYearStartMonth": 0, + "graphTooltip": 0, + "id": null, + "links": [], + "liveNow": false, + "panels": [ + { + "collapsed": false, + "gridPos": { + "h": 1, + "w": 24, + "x": 0, + "y": 0 + }, + "id": 7, + "panels": [], + "title": "Transport", + "type": "row" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 1 + }, + "id": 3, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "single", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(quicgo_server_received_packets_dropped_total{instance=~\"$instance\"}[$__rate_interval])) by (reason)", + "instant": false, + "legendFormat": "__auto", + "range": true, + "refId": "A" + } + ], + "title": "Server Dropped Packets", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 1 + }, + "id": 12, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "single", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(quicgo_server_connections_rejected_total{instance=~\"$instance\"}[$__rate_interval])) by (reason)", + "hide": true, + "instant": false, + "legendFormat": "__auto", + "range": true, + "refId": "A" + } + ], + "title": "Rejected Connections", + "type": "timeseries" + }, + { + "collapsed": false, + "gridPos": { + "h": 1, + "w": 24, + "x": 0, + "y": 9 + }, + "id": 6, + "panels": [], + "title": "Connection", + "type": "row" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "thresholds" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 10 + }, + "id": 1, + "options": { + "colorMode": "value", + "graphMode": "area", + "justifyMode": "auto", + "orientation": "auto", + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "showPercentChange": false, + "textMode": "auto", + "wideLayout": true + }, + "pluginVersion": "10.2.3", + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum (quicgo_connections_started_total{instance=~\"$instance\"}) by (dir) - sum (quicgo_connections_closed_total{instance=~\"$instance\"}) by (dir)", + "instant": false, + "legendFormat": "{{dir}}", + "range": true, + "refId": "A" + } + ], + "title": "Currently Active Connections", + "type": "stat" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + }, + "unit": "s" + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 10 + }, + "id": 5, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.5, sum(rate(quicgo_handshake_duration_seconds_bucket{instance=~\"$instance\"}[$__rate_interval])) by (le))", + "instant": false, + "legendFormat": "50th percentile", + "range": true, + "refId": "A" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.9, sum(rate(quicgo_handshake_duration_seconds_bucket{instance=~\"$instance\"}[$__rate_interval])) by (le))", + "hide": false, + "instant": false, + "legendFormat": "90th percentile", + "range": true, + "refId": "B" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.95, sum(rate(quicgo_handshake_duration_seconds_bucket{instance=~\"$instance\"}[$__rate_interval])) by (le))", + "hide": false, + "instant": false, + "legendFormat": "95th percentile", + "range": true, + "refId": "C" + } + ], + "title": "Handshake Duration", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 18 + }, + "id": 11, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "single", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(quicgo_connections_closed_total{instance=~\"$instance\"}[$__rate_interval])) by (reason)", + "instant": false, + "legendFormat": "__auto", + "range": true, + "refId": "A" + } + ], + "title": "Close Reason", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + }, + "unit": "s" + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 18 + }, + "id": 2, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.5, sum(rate(quicgo_connection_duration_seconds_bucket{instance=~\"$instance\"}[$__rate_interval])) by (le))\n", + "hide": false, + "instant": false, + "legendFormat": "50th percentile", + "range": true, + "refId": "A" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.9, sum(rate(quicgo_connection_duration_seconds_bucket{instance=~\"$instance\"}[$__rate_interval])) by (le))\n", + "hide": false, + "instant": false, + "legendFormat": "90th percentile", + "range": true, + "refId": "B" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.95, sum(rate(quicgo_connection_duration_seconds_bucket{instance=~\"$instance\"}[$__rate_interval])) by (le))\n", + "hide": false, + "instant": false, + "legendFormat": "95th percentile", + "range": true, + "refId": "C" + } + ], + "title": "Connection Durations", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 26 + }, + "id": 13, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "single", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "disableTextWrap": false, + "editorMode": "code", + "expr": "sum(rate(quicgo_packets_received_total{instance=~\"$instance\"}[$__rate_interval])) by (type)", + "fullMetaSearch": false, + "includeNullMetadata": false, + "instant": false, + "legendFormat": "{{type}}", + "range": true, + "refId": "A", + "useBackend": false + } + ], + "title": "Packets Received", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green", + "value": null + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 26 + }, + "id": 15, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "single", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "disableTextWrap": false, + "editorMode": "code", + "expr": "sum(rate(quicgo_packets_sent_total{instance=~\"$instance\"}[$__rate_interval])) by (type)", + "fullMetaSearch": false, + "includeNullMetadata": false, + "instant": false, + "legendFormat": "{{type}}", + "range": true, + "refId": "A", + "useBackend": false + } + ], + "title": "Packets Sent", + "type": "timeseries" + } + ], + "refresh": "", + "schemaVersion": 39, + "tags": [], + "templating": { + "list": [ + { + "current": {}, + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "definition": "label_values(up,instance)", + "hide": 0, + "includeAll": true, + "multi": true, + "name": "instance", + "options": [], + "query": { + "qryType": 1, + "query": "label_values(up,instance)", + "refId": "PrometheusVariableQueryEditor-VariableQuery" + }, + "refresh": 1, + "regex": "", + "skipUrlSync": false, + "sort": 0, + "type": "query" + } + ] + }, + "time": { + "from": "now-30m", + "to": "now" + }, + "timepicker": {}, + "timezone": "", + "title": "quic-go", + "uid": "afd27180-618a-42ab-99fd-0508776d9c29", + "version": 17, + "weekStart": "" +} diff --git a/third_party/quic-go/mock_ack_frame_source_test.go b/third_party/quic-go/mock_ack_frame_source_test.go new file mode 100644 index 0000000..6268a76 --- /dev/null +++ b/third_party/quic-go/mock_ack_frame_source_test.go @@ -0,0 +1,81 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: AckFrameSource) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_ack_frame_source_test.go github.com/apernet/quic-go AckFrameSource +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + wire "github.com/apernet/quic-go/internal/wire" + gomock "go.uber.org/mock/gomock" +) + +// MockAckFrameSource is a mock of AckFrameSource interface. +type MockAckFrameSource struct { + ctrl *gomock.Controller + recorder *MockAckFrameSourceMockRecorder + isgomock struct{} +} + +// MockAckFrameSourceMockRecorder is the mock recorder for MockAckFrameSource. +type MockAckFrameSourceMockRecorder struct { + mock *MockAckFrameSource +} + +// NewMockAckFrameSource creates a new mock instance. +func NewMockAckFrameSource(ctrl *gomock.Controller) *MockAckFrameSource { + mock := &MockAckFrameSource{ctrl: ctrl} + mock.recorder = &MockAckFrameSourceMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockAckFrameSource) EXPECT() *MockAckFrameSourceMockRecorder { + return m.recorder +} + +// GetAckFrame mocks base method. +func (m *MockAckFrameSource) GetAckFrame(arg0 protocol.EncryptionLevel, now monotime.Time, onlyIfQueued bool) *wire.AckFrame { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAckFrame", arg0, now, onlyIfQueued) + ret0, _ := ret[0].(*wire.AckFrame) + return ret0 +} + +// GetAckFrame indicates an expected call of GetAckFrame. +func (mr *MockAckFrameSourceMockRecorder) GetAckFrame(arg0, now, onlyIfQueued any) *MockAckFrameSourceGetAckFrameCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAckFrame", reflect.TypeOf((*MockAckFrameSource)(nil).GetAckFrame), arg0, now, onlyIfQueued) + return &MockAckFrameSourceGetAckFrameCall{Call: call} +} + +// MockAckFrameSourceGetAckFrameCall wrap *gomock.Call +type MockAckFrameSourceGetAckFrameCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockAckFrameSourceGetAckFrameCall) Return(arg0 *wire.AckFrame) *MockAckFrameSourceGetAckFrameCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockAckFrameSourceGetAckFrameCall) Do(f func(protocol.EncryptionLevel, monotime.Time, bool) *wire.AckFrame) *MockAckFrameSourceGetAckFrameCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockAckFrameSourceGetAckFrameCall) DoAndReturn(f func(protocol.EncryptionLevel, monotime.Time, bool) *wire.AckFrame) *MockAckFrameSourceGetAckFrameCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_conn_runner_test.go b/third_party/quic-go/mock_conn_runner_test.go new file mode 100644 index 0000000..76af294 --- /dev/null +++ b/third_party/quic-go/mock_conn_runner_test.go @@ -0,0 +1,224 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: ConnRunner) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_conn_runner_test.go github.com/apernet/quic-go ConnRunner +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + time "time" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockConnRunner is a mock of ConnRunner interface. +type MockConnRunner struct { + ctrl *gomock.Controller + recorder *MockConnRunnerMockRecorder + isgomock struct{} +} + +// MockConnRunnerMockRecorder is the mock recorder for MockConnRunner. +type MockConnRunnerMockRecorder struct { + mock *MockConnRunner +} + +// NewMockConnRunner creates a new mock instance. +func NewMockConnRunner(ctrl *gomock.Controller) *MockConnRunner { + mock := &MockConnRunner{ctrl: ctrl} + mock.recorder = &MockConnRunnerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockConnRunner) EXPECT() *MockConnRunnerMockRecorder { + return m.recorder +} + +// Add mocks base method. +func (m *MockConnRunner) Add(arg0 protocol.ConnectionID, arg1 packetHandler) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Add", arg0, arg1) + ret0, _ := ret[0].(bool) + return ret0 +} + +// Add indicates an expected call of Add. +func (mr *MockConnRunnerMockRecorder) Add(arg0, arg1 any) *MockConnRunnerAddCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Add", reflect.TypeOf((*MockConnRunner)(nil).Add), arg0, arg1) + return &MockConnRunnerAddCall{Call: call} +} + +// MockConnRunnerAddCall wrap *gomock.Call +type MockConnRunnerAddCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockConnRunnerAddCall) Return(arg0 bool) *MockConnRunnerAddCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockConnRunnerAddCall) Do(f func(protocol.ConnectionID, packetHandler) bool) *MockConnRunnerAddCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockConnRunnerAddCall) DoAndReturn(f func(protocol.ConnectionID, packetHandler) bool) *MockConnRunnerAddCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// AddResetToken mocks base method. +func (m *MockConnRunner) AddResetToken(arg0 protocol.StatelessResetToken, arg1 packetHandler) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "AddResetToken", arg0, arg1) +} + +// AddResetToken indicates an expected call of AddResetToken. +func (mr *MockConnRunnerMockRecorder) AddResetToken(arg0, arg1 any) *MockConnRunnerAddResetTokenCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AddResetToken", reflect.TypeOf((*MockConnRunner)(nil).AddResetToken), arg0, arg1) + return &MockConnRunnerAddResetTokenCall{Call: call} +} + +// MockConnRunnerAddResetTokenCall wrap *gomock.Call +type MockConnRunnerAddResetTokenCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockConnRunnerAddResetTokenCall) Return() *MockConnRunnerAddResetTokenCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockConnRunnerAddResetTokenCall) Do(f func(protocol.StatelessResetToken, packetHandler)) *MockConnRunnerAddResetTokenCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockConnRunnerAddResetTokenCall) DoAndReturn(f func(protocol.StatelessResetToken, packetHandler)) *MockConnRunnerAddResetTokenCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Remove mocks base method. +func (m *MockConnRunner) Remove(arg0 protocol.ConnectionID) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Remove", arg0) +} + +// Remove indicates an expected call of Remove. +func (mr *MockConnRunnerMockRecorder) Remove(arg0 any) *MockConnRunnerRemoveCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Remove", reflect.TypeOf((*MockConnRunner)(nil).Remove), arg0) + return &MockConnRunnerRemoveCall{Call: call} +} + +// MockConnRunnerRemoveCall wrap *gomock.Call +type MockConnRunnerRemoveCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockConnRunnerRemoveCall) Return() *MockConnRunnerRemoveCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockConnRunnerRemoveCall) Do(f func(protocol.ConnectionID)) *MockConnRunnerRemoveCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockConnRunnerRemoveCall) DoAndReturn(f func(protocol.ConnectionID)) *MockConnRunnerRemoveCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// RemoveResetToken mocks base method. +func (m *MockConnRunner) RemoveResetToken(arg0 protocol.StatelessResetToken) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "RemoveResetToken", arg0) +} + +// RemoveResetToken indicates an expected call of RemoveResetToken. +func (mr *MockConnRunnerMockRecorder) RemoveResetToken(arg0 any) *MockConnRunnerRemoveResetTokenCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RemoveResetToken", reflect.TypeOf((*MockConnRunner)(nil).RemoveResetToken), arg0) + return &MockConnRunnerRemoveResetTokenCall{Call: call} +} + +// MockConnRunnerRemoveResetTokenCall wrap *gomock.Call +type MockConnRunnerRemoveResetTokenCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockConnRunnerRemoveResetTokenCall) Return() *MockConnRunnerRemoveResetTokenCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockConnRunnerRemoveResetTokenCall) Do(f func(protocol.StatelessResetToken)) *MockConnRunnerRemoveResetTokenCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockConnRunnerRemoveResetTokenCall) DoAndReturn(f func(protocol.StatelessResetToken)) *MockConnRunnerRemoveResetTokenCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReplaceWithClosed mocks base method. +func (m *MockConnRunner) ReplaceWithClosed(arg0 []protocol.ConnectionID, arg1 []byte, arg2 time.Duration) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "ReplaceWithClosed", arg0, arg1, arg2) +} + +// ReplaceWithClosed indicates an expected call of ReplaceWithClosed. +func (mr *MockConnRunnerMockRecorder) ReplaceWithClosed(arg0, arg1, arg2 any) *MockConnRunnerReplaceWithClosedCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReplaceWithClosed", reflect.TypeOf((*MockConnRunner)(nil).ReplaceWithClosed), arg0, arg1, arg2) + return &MockConnRunnerReplaceWithClosedCall{Call: call} +} + +// MockConnRunnerReplaceWithClosedCall wrap *gomock.Call +type MockConnRunnerReplaceWithClosedCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockConnRunnerReplaceWithClosedCall) Return() *MockConnRunnerReplaceWithClosedCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockConnRunnerReplaceWithClosedCall) Do(f func([]protocol.ConnectionID, []byte, time.Duration)) *MockConnRunnerReplaceWithClosedCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockConnRunnerReplaceWithClosedCall) DoAndReturn(f func([]protocol.ConnectionID, []byte, time.Duration)) *MockConnRunnerReplaceWithClosedCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_frame_source_test.go b/third_party/quic-go/mock_frame_source_test.go new file mode 100644 index 0000000..c12dba1 --- /dev/null +++ b/third_party/quic-go/mock_frame_source_test.go @@ -0,0 +1,121 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: FrameSource) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_frame_source_test.go github.com/apernet/quic-go FrameSource +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + ackhandler "github.com/apernet/quic-go/internal/ackhandler" + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockFrameSource is a mock of FrameSource interface. +type MockFrameSource struct { + ctrl *gomock.Controller + recorder *MockFrameSourceMockRecorder + isgomock struct{} +} + +// MockFrameSourceMockRecorder is the mock recorder for MockFrameSource. +type MockFrameSourceMockRecorder struct { + mock *MockFrameSource +} + +// NewMockFrameSource creates a new mock instance. +func NewMockFrameSource(ctrl *gomock.Controller) *MockFrameSource { + mock := &MockFrameSource{ctrl: ctrl} + mock.recorder = &MockFrameSourceMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockFrameSource) EXPECT() *MockFrameSourceMockRecorder { + return m.recorder +} + +// Append mocks base method. +func (m *MockFrameSource) Append(arg0 []ackhandler.Frame, arg1 []ackhandler.StreamFrame, arg2 protocol.ByteCount, arg3 monotime.Time, arg4 protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Append", arg0, arg1, arg2, arg3, arg4) + ret0, _ := ret[0].([]ackhandler.Frame) + ret1, _ := ret[1].([]ackhandler.StreamFrame) + ret2, _ := ret[2].(protocol.ByteCount) + return ret0, ret1, ret2 +} + +// Append indicates an expected call of Append. +func (mr *MockFrameSourceMockRecorder) Append(arg0, arg1, arg2, arg3, arg4 any) *MockFrameSourceAppendCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Append", reflect.TypeOf((*MockFrameSource)(nil).Append), arg0, arg1, arg2, arg3, arg4) + return &MockFrameSourceAppendCall{Call: call} +} + +// MockFrameSourceAppendCall wrap *gomock.Call +type MockFrameSourceAppendCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockFrameSourceAppendCall) Return(arg0 []ackhandler.Frame, arg1 []ackhandler.StreamFrame, arg2 protocol.ByteCount) *MockFrameSourceAppendCall { + c.Call = c.Call.Return(arg0, arg1, arg2) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockFrameSourceAppendCall) Do(f func([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount, monotime.Time, protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount)) *MockFrameSourceAppendCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockFrameSourceAppendCall) DoAndReturn(f func([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount, monotime.Time, protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount)) *MockFrameSourceAppendCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// HasData mocks base method. +func (m *MockFrameSource) HasData() bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "HasData") + ret0, _ := ret[0].(bool) + return ret0 +} + +// HasData indicates an expected call of HasData. +func (mr *MockFrameSourceMockRecorder) HasData() *MockFrameSourceHasDataCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasData", reflect.TypeOf((*MockFrameSource)(nil).HasData)) + return &MockFrameSourceHasDataCall{Call: call} +} + +// MockFrameSourceHasDataCall wrap *gomock.Call +type MockFrameSourceHasDataCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockFrameSourceHasDataCall) Return(arg0 bool) *MockFrameSourceHasDataCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockFrameSourceHasDataCall) Do(f func() bool) *MockFrameSourceHasDataCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockFrameSourceHasDataCall) DoAndReturn(f func() bool) *MockFrameSourceHasDataCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_mtu_discoverer_test.go b/third_party/quic-go/mock_mtu_discoverer_test.go new file mode 100644 index 0000000..8b9ff12 --- /dev/null +++ b/third_party/quic-go/mock_mtu_discoverer_test.go @@ -0,0 +1,230 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: MTUDiscoverer) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_mtu_discoverer_test.go github.com/apernet/quic-go MTUDiscoverer +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + ackhandler "github.com/apernet/quic-go/internal/ackhandler" + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockMTUDiscoverer is a mock of MTUDiscoverer interface. +type MockMTUDiscoverer struct { + ctrl *gomock.Controller + recorder *MockMTUDiscovererMockRecorder + isgomock struct{} +} + +// MockMTUDiscovererMockRecorder is the mock recorder for MockMTUDiscoverer. +type MockMTUDiscovererMockRecorder struct { + mock *MockMTUDiscoverer +} + +// NewMockMTUDiscoverer creates a new mock instance. +func NewMockMTUDiscoverer(ctrl *gomock.Controller) *MockMTUDiscoverer { + mock := &MockMTUDiscoverer{ctrl: ctrl} + mock.recorder = &MockMTUDiscovererMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockMTUDiscoverer) EXPECT() *MockMTUDiscovererMockRecorder { + return m.recorder +} + +// CurrentSize mocks base method. +func (m *MockMTUDiscoverer) CurrentSize() protocol.ByteCount { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "CurrentSize") + ret0, _ := ret[0].(protocol.ByteCount) + return ret0 +} + +// CurrentSize indicates an expected call of CurrentSize. +func (mr *MockMTUDiscovererMockRecorder) CurrentSize() *MockMTUDiscovererCurrentSizeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CurrentSize", reflect.TypeOf((*MockMTUDiscoverer)(nil).CurrentSize)) + return &MockMTUDiscovererCurrentSizeCall{Call: call} +} + +// MockMTUDiscovererCurrentSizeCall wrap *gomock.Call +type MockMTUDiscovererCurrentSizeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockMTUDiscovererCurrentSizeCall) Return(arg0 protocol.ByteCount) *MockMTUDiscovererCurrentSizeCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockMTUDiscovererCurrentSizeCall) Do(f func() protocol.ByteCount) *MockMTUDiscovererCurrentSizeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockMTUDiscovererCurrentSizeCall) DoAndReturn(f func() protocol.ByteCount) *MockMTUDiscovererCurrentSizeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetPing mocks base method. +func (m *MockMTUDiscoverer) GetPing(now monotime.Time) (ackhandler.Frame, protocol.ByteCount) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPing", now) + ret0, _ := ret[0].(ackhandler.Frame) + ret1, _ := ret[1].(protocol.ByteCount) + return ret0, ret1 +} + +// GetPing indicates an expected call of GetPing. +func (mr *MockMTUDiscovererMockRecorder) GetPing(now any) *MockMTUDiscovererGetPingCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPing", reflect.TypeOf((*MockMTUDiscoverer)(nil).GetPing), now) + return &MockMTUDiscovererGetPingCall{Call: call} +} + +// MockMTUDiscovererGetPingCall wrap *gomock.Call +type MockMTUDiscovererGetPingCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockMTUDiscovererGetPingCall) Return(ping ackhandler.Frame, datagramSize protocol.ByteCount) *MockMTUDiscovererGetPingCall { + c.Call = c.Call.Return(ping, datagramSize) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockMTUDiscovererGetPingCall) Do(f func(monotime.Time) (ackhandler.Frame, protocol.ByteCount)) *MockMTUDiscovererGetPingCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockMTUDiscovererGetPingCall) DoAndReturn(f func(monotime.Time) (ackhandler.Frame, protocol.ByteCount)) *MockMTUDiscovererGetPingCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Reset mocks base method. +func (m *MockMTUDiscoverer) Reset(now monotime.Time, start, max protocol.ByteCount) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Reset", now, start, max) +} + +// Reset indicates an expected call of Reset. +func (mr *MockMTUDiscovererMockRecorder) Reset(now, start, max any) *MockMTUDiscovererResetCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Reset", reflect.TypeOf((*MockMTUDiscoverer)(nil).Reset), now, start, max) + return &MockMTUDiscovererResetCall{Call: call} +} + +// MockMTUDiscovererResetCall wrap *gomock.Call +type MockMTUDiscovererResetCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockMTUDiscovererResetCall) Return() *MockMTUDiscovererResetCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockMTUDiscovererResetCall) Do(f func(monotime.Time, protocol.ByteCount, protocol.ByteCount)) *MockMTUDiscovererResetCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockMTUDiscovererResetCall) DoAndReturn(f func(monotime.Time, protocol.ByteCount, protocol.ByteCount)) *MockMTUDiscovererResetCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ShouldSendProbe mocks base method. +func (m *MockMTUDiscoverer) ShouldSendProbe(now monotime.Time) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ShouldSendProbe", now) + ret0, _ := ret[0].(bool) + return ret0 +} + +// ShouldSendProbe indicates an expected call of ShouldSendProbe. +func (mr *MockMTUDiscovererMockRecorder) ShouldSendProbe(now any) *MockMTUDiscovererShouldSendProbeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ShouldSendProbe", reflect.TypeOf((*MockMTUDiscoverer)(nil).ShouldSendProbe), now) + return &MockMTUDiscovererShouldSendProbeCall{Call: call} +} + +// MockMTUDiscovererShouldSendProbeCall wrap *gomock.Call +type MockMTUDiscovererShouldSendProbeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockMTUDiscovererShouldSendProbeCall) Return(arg0 bool) *MockMTUDiscovererShouldSendProbeCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockMTUDiscovererShouldSendProbeCall) Do(f func(monotime.Time) bool) *MockMTUDiscovererShouldSendProbeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockMTUDiscovererShouldSendProbeCall) DoAndReturn(f func(monotime.Time) bool) *MockMTUDiscovererShouldSendProbeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Start mocks base method. +func (m *MockMTUDiscoverer) Start(now monotime.Time) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Start", now) +} + +// Start indicates an expected call of Start. +func (mr *MockMTUDiscovererMockRecorder) Start(now any) *MockMTUDiscovererStartCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Start", reflect.TypeOf((*MockMTUDiscoverer)(nil).Start), now) + return &MockMTUDiscovererStartCall{Call: call} +} + +// MockMTUDiscovererStartCall wrap *gomock.Call +type MockMTUDiscovererStartCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockMTUDiscovererStartCall) Return() *MockMTUDiscovererStartCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockMTUDiscovererStartCall) Do(f func(monotime.Time)) *MockMTUDiscovererStartCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockMTUDiscovererStartCall) DoAndReturn(f func(monotime.Time)) *MockMTUDiscovererStartCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_packer_test.go b/third_party/quic-go/mock_packer_test.go new file mode 100644 index 0000000..51ab2c2 --- /dev/null +++ b/third_party/quic-go/mock_packer_test.go @@ -0,0 +1,395 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: Packer) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_packer_test.go github.com/apernet/quic-go Packer +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + ackhandler "github.com/apernet/quic-go/internal/ackhandler" + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + qerr "github.com/apernet/quic-go/internal/qerr" + gomock "go.uber.org/mock/gomock" +) + +// MockPacker is a mock of Packer interface. +type MockPacker struct { + ctrl *gomock.Controller + recorder *MockPackerMockRecorder + isgomock struct{} +} + +// MockPackerMockRecorder is the mock recorder for MockPacker. +type MockPackerMockRecorder struct { + mock *MockPacker +} + +// NewMockPacker creates a new mock instance. +func NewMockPacker(ctrl *gomock.Controller) *MockPacker { + mock := &MockPacker{ctrl: ctrl} + mock.recorder = &MockPackerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockPacker) EXPECT() *MockPackerMockRecorder { + return m.recorder +} + +// AppendPacket mocks base method. +func (m *MockPacker) AppendPacket(arg0 *packetBuffer, maxPacketSize protocol.ByteCount, now monotime.Time, v protocol.Version) (shortHeaderPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "AppendPacket", arg0, maxPacketSize, now, v) + ret0, _ := ret[0].(shortHeaderPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// AppendPacket indicates an expected call of AppendPacket. +func (mr *MockPackerMockRecorder) AppendPacket(arg0, maxPacketSize, now, v any) *MockPackerAppendPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AppendPacket", reflect.TypeOf((*MockPacker)(nil).AppendPacket), arg0, maxPacketSize, now, v) + return &MockPackerAppendPacketCall{Call: call} +} + +// MockPackerAppendPacketCall wrap *gomock.Call +type MockPackerAppendPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerAppendPacketCall) Return(arg0 shortHeaderPacket, arg1 error) *MockPackerAppendPacketCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerAppendPacketCall) Do(f func(*packetBuffer, protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, error)) *MockPackerAppendPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerAppendPacketCall) DoAndReturn(f func(*packetBuffer, protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, error)) *MockPackerAppendPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackAckOnlyPacket mocks base method. +func (m *MockPacker) PackAckOnlyPacket(maxPacketSize protocol.ByteCount, now monotime.Time, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackAckOnlyPacket", maxPacketSize, now, v) + ret0, _ := ret[0].(shortHeaderPacket) + ret1, _ := ret[1].(*packetBuffer) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// PackAckOnlyPacket indicates an expected call of PackAckOnlyPacket. +func (mr *MockPackerMockRecorder) PackAckOnlyPacket(maxPacketSize, now, v any) *MockPackerPackAckOnlyPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackAckOnlyPacket", reflect.TypeOf((*MockPacker)(nil).PackAckOnlyPacket), maxPacketSize, now, v) + return &MockPackerPackAckOnlyPacketCall{Call: call} +} + +// MockPackerPackAckOnlyPacketCall wrap *gomock.Call +type MockPackerPackAckOnlyPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackAckOnlyPacketCall) Return(arg0 shortHeaderPacket, arg1 *packetBuffer, arg2 error) *MockPackerPackAckOnlyPacketCall { + c.Call = c.Call.Return(arg0, arg1, arg2) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackAckOnlyPacketCall) Do(f func(protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, *packetBuffer, error)) *MockPackerPackAckOnlyPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackAckOnlyPacketCall) DoAndReturn(f func(protocol.ByteCount, monotime.Time, protocol.Version) (shortHeaderPacket, *packetBuffer, error)) *MockPackerPackAckOnlyPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackApplicationClose mocks base method. +func (m *MockPacker) PackApplicationClose(arg0 *qerr.ApplicationError, arg1 protocol.ByteCount, arg2 protocol.Version) (*coalescedPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackApplicationClose", arg0, arg1, arg2) + ret0, _ := ret[0].(*coalescedPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// PackApplicationClose indicates an expected call of PackApplicationClose. +func (mr *MockPackerMockRecorder) PackApplicationClose(arg0, arg1, arg2 any) *MockPackerPackApplicationCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackApplicationClose", reflect.TypeOf((*MockPacker)(nil).PackApplicationClose), arg0, arg1, arg2) + return &MockPackerPackApplicationCloseCall{Call: call} +} + +// MockPackerPackApplicationCloseCall wrap *gomock.Call +type MockPackerPackApplicationCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackApplicationCloseCall) Return(arg0 *coalescedPacket, arg1 error) *MockPackerPackApplicationCloseCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackApplicationCloseCall) Do(f func(*qerr.ApplicationError, protocol.ByteCount, protocol.Version) (*coalescedPacket, error)) *MockPackerPackApplicationCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackApplicationCloseCall) DoAndReturn(f func(*qerr.ApplicationError, protocol.ByteCount, protocol.Version) (*coalescedPacket, error)) *MockPackerPackApplicationCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackCoalescedPacket mocks base method. +func (m *MockPacker) PackCoalescedPacket(onlyAck bool, maxPacketSize protocol.ByteCount, now monotime.Time, v protocol.Version) (*coalescedPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackCoalescedPacket", onlyAck, maxPacketSize, now, v) + ret0, _ := ret[0].(*coalescedPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// PackCoalescedPacket indicates an expected call of PackCoalescedPacket. +func (mr *MockPackerMockRecorder) PackCoalescedPacket(onlyAck, maxPacketSize, now, v any) *MockPackerPackCoalescedPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackCoalescedPacket", reflect.TypeOf((*MockPacker)(nil).PackCoalescedPacket), onlyAck, maxPacketSize, now, v) + return &MockPackerPackCoalescedPacketCall{Call: call} +} + +// MockPackerPackCoalescedPacketCall wrap *gomock.Call +type MockPackerPackCoalescedPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackCoalescedPacketCall) Return(arg0 *coalescedPacket, arg1 error) *MockPackerPackCoalescedPacketCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackCoalescedPacketCall) Do(f func(bool, protocol.ByteCount, monotime.Time, protocol.Version) (*coalescedPacket, error)) *MockPackerPackCoalescedPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackCoalescedPacketCall) DoAndReturn(f func(bool, protocol.ByteCount, monotime.Time, protocol.Version) (*coalescedPacket, error)) *MockPackerPackCoalescedPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackConnectionClose mocks base method. +func (m *MockPacker) PackConnectionClose(arg0 *qerr.TransportError, arg1 protocol.ByteCount, arg2 protocol.Version) (*coalescedPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackConnectionClose", arg0, arg1, arg2) + ret0, _ := ret[0].(*coalescedPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// PackConnectionClose indicates an expected call of PackConnectionClose. +func (mr *MockPackerMockRecorder) PackConnectionClose(arg0, arg1, arg2 any) *MockPackerPackConnectionCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackConnectionClose", reflect.TypeOf((*MockPacker)(nil).PackConnectionClose), arg0, arg1, arg2) + return &MockPackerPackConnectionCloseCall{Call: call} +} + +// MockPackerPackConnectionCloseCall wrap *gomock.Call +type MockPackerPackConnectionCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackConnectionCloseCall) Return(arg0 *coalescedPacket, arg1 error) *MockPackerPackConnectionCloseCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackConnectionCloseCall) Do(f func(*qerr.TransportError, protocol.ByteCount, protocol.Version) (*coalescedPacket, error)) *MockPackerPackConnectionCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackConnectionCloseCall) DoAndReturn(f func(*qerr.TransportError, protocol.ByteCount, protocol.Version) (*coalescedPacket, error)) *MockPackerPackConnectionCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackMTUProbePacket mocks base method. +func (m *MockPacker) PackMTUProbePacket(ping ackhandler.Frame, size protocol.ByteCount, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackMTUProbePacket", ping, size, v) + ret0, _ := ret[0].(shortHeaderPacket) + ret1, _ := ret[1].(*packetBuffer) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// PackMTUProbePacket indicates an expected call of PackMTUProbePacket. +func (mr *MockPackerMockRecorder) PackMTUProbePacket(ping, size, v any) *MockPackerPackMTUProbePacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackMTUProbePacket", reflect.TypeOf((*MockPacker)(nil).PackMTUProbePacket), ping, size, v) + return &MockPackerPackMTUProbePacketCall{Call: call} +} + +// MockPackerPackMTUProbePacketCall wrap *gomock.Call +type MockPackerPackMTUProbePacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackMTUProbePacketCall) Return(arg0 shortHeaderPacket, arg1 *packetBuffer, arg2 error) *MockPackerPackMTUProbePacketCall { + c.Call = c.Call.Return(arg0, arg1, arg2) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackMTUProbePacketCall) Do(f func(ackhandler.Frame, protocol.ByteCount, protocol.Version) (shortHeaderPacket, *packetBuffer, error)) *MockPackerPackMTUProbePacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackMTUProbePacketCall) DoAndReturn(f func(ackhandler.Frame, protocol.ByteCount, protocol.Version) (shortHeaderPacket, *packetBuffer, error)) *MockPackerPackMTUProbePacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackPTOProbePacket mocks base method. +func (m *MockPacker) PackPTOProbePacket(arg0 protocol.EncryptionLevel, arg1 protocol.ByteCount, addPingIfEmpty bool, now monotime.Time, v protocol.Version) (*coalescedPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackPTOProbePacket", arg0, arg1, addPingIfEmpty, now, v) + ret0, _ := ret[0].(*coalescedPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// PackPTOProbePacket indicates an expected call of PackPTOProbePacket. +func (mr *MockPackerMockRecorder) PackPTOProbePacket(arg0, arg1, addPingIfEmpty, now, v any) *MockPackerPackPTOProbePacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackPTOProbePacket", reflect.TypeOf((*MockPacker)(nil).PackPTOProbePacket), arg0, arg1, addPingIfEmpty, now, v) + return &MockPackerPackPTOProbePacketCall{Call: call} +} + +// MockPackerPackPTOProbePacketCall wrap *gomock.Call +type MockPackerPackPTOProbePacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackPTOProbePacketCall) Return(arg0 *coalescedPacket, arg1 error) *MockPackerPackPTOProbePacketCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackPTOProbePacketCall) Do(f func(protocol.EncryptionLevel, protocol.ByteCount, bool, monotime.Time, protocol.Version) (*coalescedPacket, error)) *MockPackerPackPTOProbePacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackPTOProbePacketCall) DoAndReturn(f func(protocol.EncryptionLevel, protocol.ByteCount, bool, monotime.Time, protocol.Version) (*coalescedPacket, error)) *MockPackerPackPTOProbePacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// PackPathProbePacket mocks base method. +func (m *MockPacker) PackPathProbePacket(arg0 protocol.ConnectionID, arg1 []ackhandler.Frame, arg2 protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "PackPathProbePacket", arg0, arg1, arg2) + ret0, _ := ret[0].(shortHeaderPacket) + ret1, _ := ret[1].(*packetBuffer) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// PackPathProbePacket indicates an expected call of PackPathProbePacket. +func (mr *MockPackerMockRecorder) PackPathProbePacket(arg0, arg1, arg2 any) *MockPackerPackPathProbePacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PackPathProbePacket", reflect.TypeOf((*MockPacker)(nil).PackPathProbePacket), arg0, arg1, arg2) + return &MockPackerPackPathProbePacketCall{Call: call} +} + +// MockPackerPackPathProbePacketCall wrap *gomock.Call +type MockPackerPackPathProbePacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerPackPathProbePacketCall) Return(arg0 shortHeaderPacket, arg1 *packetBuffer, arg2 error) *MockPackerPackPathProbePacketCall { + c.Call = c.Call.Return(arg0, arg1, arg2) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerPackPathProbePacketCall) Do(f func(protocol.ConnectionID, []ackhandler.Frame, protocol.Version) (shortHeaderPacket, *packetBuffer, error)) *MockPackerPackPathProbePacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerPackPathProbePacketCall) DoAndReturn(f func(protocol.ConnectionID, []ackhandler.Frame, protocol.Version) (shortHeaderPacket, *packetBuffer, error)) *MockPackerPackPathProbePacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetToken mocks base method. +func (m *MockPacker) SetToken(arg0 []byte) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetToken", arg0) +} + +// SetToken indicates an expected call of SetToken. +func (mr *MockPackerMockRecorder) SetToken(arg0 any) *MockPackerSetTokenCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetToken", reflect.TypeOf((*MockPacker)(nil).SetToken), arg0) + return &MockPackerSetTokenCall{Call: call} +} + +// MockPackerSetTokenCall wrap *gomock.Call +type MockPackerSetTokenCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPackerSetTokenCall) Return() *MockPackerSetTokenCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPackerSetTokenCall) Do(f func([]byte)) *MockPackerSetTokenCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPackerSetTokenCall) DoAndReturn(f func([]byte)) *MockPackerSetTokenCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_packet_handler_test.go b/third_party/quic-go/mock_packet_handler_test.go new file mode 100644 index 0000000..3c99a8e --- /dev/null +++ b/third_party/quic-go/mock_packet_handler_test.go @@ -0,0 +1,149 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: PacketHandler) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_packet_handler_test.go github.com/apernet/quic-go PacketHandler +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + qerr "github.com/apernet/quic-go/internal/qerr" + gomock "go.uber.org/mock/gomock" +) + +// MockPacketHandler is a mock of PacketHandler interface. +type MockPacketHandler struct { + ctrl *gomock.Controller + recorder *MockPacketHandlerMockRecorder + isgomock struct{} +} + +// MockPacketHandlerMockRecorder is the mock recorder for MockPacketHandler. +type MockPacketHandlerMockRecorder struct { + mock *MockPacketHandler +} + +// NewMockPacketHandler creates a new mock instance. +func NewMockPacketHandler(ctrl *gomock.Controller) *MockPacketHandler { + mock := &MockPacketHandler{ctrl: ctrl} + mock.recorder = &MockPacketHandlerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockPacketHandler) EXPECT() *MockPacketHandlerMockRecorder { + return m.recorder +} + +// closeWithTransportError mocks base method. +func (m *MockPacketHandler) closeWithTransportError(arg0 qerr.TransportErrorCode) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "closeWithTransportError", arg0) +} + +// closeWithTransportError indicates an expected call of closeWithTransportError. +func (mr *MockPacketHandlerMockRecorder) closeWithTransportError(arg0 any) *MockPacketHandlercloseWithTransportErrorCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "closeWithTransportError", reflect.TypeOf((*MockPacketHandler)(nil).closeWithTransportError), arg0) + return &MockPacketHandlercloseWithTransportErrorCall{Call: call} +} + +// MockPacketHandlercloseWithTransportErrorCall wrap *gomock.Call +type MockPacketHandlercloseWithTransportErrorCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketHandlercloseWithTransportErrorCall) Return() *MockPacketHandlercloseWithTransportErrorCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketHandlercloseWithTransportErrorCall) Do(f func(qerr.TransportErrorCode)) *MockPacketHandlercloseWithTransportErrorCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketHandlercloseWithTransportErrorCall) DoAndReturn(f func(qerr.TransportErrorCode)) *MockPacketHandlercloseWithTransportErrorCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// destroy mocks base method. +func (m *MockPacketHandler) destroy(arg0 error) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "destroy", arg0) +} + +// destroy indicates an expected call of destroy. +func (mr *MockPacketHandlerMockRecorder) destroy(arg0 any) *MockPacketHandlerdestroyCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "destroy", reflect.TypeOf((*MockPacketHandler)(nil).destroy), arg0) + return &MockPacketHandlerdestroyCall{Call: call} +} + +// MockPacketHandlerdestroyCall wrap *gomock.Call +type MockPacketHandlerdestroyCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketHandlerdestroyCall) Return() *MockPacketHandlerdestroyCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketHandlerdestroyCall) Do(f func(error)) *MockPacketHandlerdestroyCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketHandlerdestroyCall) DoAndReturn(f func(error)) *MockPacketHandlerdestroyCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// handlePacket mocks base method. +func (m *MockPacketHandler) handlePacket(arg0 receivedPacket) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "handlePacket", arg0) +} + +// handlePacket indicates an expected call of handlePacket. +func (mr *MockPacketHandlerMockRecorder) handlePacket(arg0 any) *MockPacketHandlerhandlePacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "handlePacket", reflect.TypeOf((*MockPacketHandler)(nil).handlePacket), arg0) + return &MockPacketHandlerhandlePacketCall{Call: call} +} + +// MockPacketHandlerhandlePacketCall wrap *gomock.Call +type MockPacketHandlerhandlePacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketHandlerhandlePacketCall) Return() *MockPacketHandlerhandlePacketCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketHandlerhandlePacketCall) Do(f func(receivedPacket)) *MockPacketHandlerhandlePacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketHandlerhandlePacketCall) DoAndReturn(f func(receivedPacket)) *MockPacketHandlerhandlePacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_packetconn_test.go b/third_party/quic-go/mock_packetconn_test.go new file mode 100644 index 0000000..9030b16 --- /dev/null +++ b/third_party/quic-go/mock_packetconn_test.go @@ -0,0 +1,311 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: net (interfaces: PacketConn) +// +// Generated by this command: +// +// mockgen -typed -package quic -self_package github.com/apernet/quic-go -self_package github.com/apernet/quic-go -destination mock_packetconn_test.go net PacketConn +// + +// Package quic is a generated GoMock package. +package quic + +import ( + net "net" + reflect "reflect" + time "time" + + gomock "go.uber.org/mock/gomock" +) + +// MockPacketConn is a mock of PacketConn interface. +type MockPacketConn struct { + ctrl *gomock.Controller + recorder *MockPacketConnMockRecorder + isgomock struct{} +} + +// MockPacketConnMockRecorder is the mock recorder for MockPacketConn. +type MockPacketConnMockRecorder struct { + mock *MockPacketConn +} + +// NewMockPacketConn creates a new mock instance. +func NewMockPacketConn(ctrl *gomock.Controller) *MockPacketConn { + mock := &MockPacketConn{ctrl: ctrl} + mock.recorder = &MockPacketConnMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockPacketConn) EXPECT() *MockPacketConnMockRecorder { + return m.recorder +} + +// Close mocks base method. +func (m *MockPacketConn) Close() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Close") + ret0, _ := ret[0].(error) + return ret0 +} + +// Close indicates an expected call of Close. +func (mr *MockPacketConnMockRecorder) Close() *MockPacketConnCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockPacketConn)(nil).Close)) + return &MockPacketConnCloseCall{Call: call} +} + +// MockPacketConnCloseCall wrap *gomock.Call +type MockPacketConnCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnCloseCall) Return(arg0 error) *MockPacketConnCloseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnCloseCall) Do(f func() error) *MockPacketConnCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnCloseCall) DoAndReturn(f func() error) *MockPacketConnCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// LocalAddr mocks base method. +func (m *MockPacketConn) LocalAddr() net.Addr { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "LocalAddr") + ret0, _ := ret[0].(net.Addr) + return ret0 +} + +// LocalAddr indicates an expected call of LocalAddr. +func (mr *MockPacketConnMockRecorder) LocalAddr() *MockPacketConnLocalAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LocalAddr", reflect.TypeOf((*MockPacketConn)(nil).LocalAddr)) + return &MockPacketConnLocalAddrCall{Call: call} +} + +// MockPacketConnLocalAddrCall wrap *gomock.Call +type MockPacketConnLocalAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnLocalAddrCall) Return(arg0 net.Addr) *MockPacketConnLocalAddrCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnLocalAddrCall) Do(f func() net.Addr) *MockPacketConnLocalAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnLocalAddrCall) DoAndReturn(f func() net.Addr) *MockPacketConnLocalAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReadFrom mocks base method. +func (m *MockPacketConn) ReadFrom(p []byte) (int, net.Addr, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReadFrom", p) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(net.Addr) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// ReadFrom indicates an expected call of ReadFrom. +func (mr *MockPacketConnMockRecorder) ReadFrom(p any) *MockPacketConnReadFromCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReadFrom", reflect.TypeOf((*MockPacketConn)(nil).ReadFrom), p) + return &MockPacketConnReadFromCall{Call: call} +} + +// MockPacketConnReadFromCall wrap *gomock.Call +type MockPacketConnReadFromCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnReadFromCall) Return(n int, addr net.Addr, err error) *MockPacketConnReadFromCall { + c.Call = c.Call.Return(n, addr, err) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnReadFromCall) Do(f func([]byte) (int, net.Addr, error)) *MockPacketConnReadFromCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnReadFromCall) DoAndReturn(f func([]byte) (int, net.Addr, error)) *MockPacketConnReadFromCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetDeadline mocks base method. +func (m *MockPacketConn) SetDeadline(t time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetDeadline", t) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetDeadline indicates an expected call of SetDeadline. +func (mr *MockPacketConnMockRecorder) SetDeadline(t any) *MockPacketConnSetDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetDeadline", reflect.TypeOf((*MockPacketConn)(nil).SetDeadline), t) + return &MockPacketConnSetDeadlineCall{Call: call} +} + +// MockPacketConnSetDeadlineCall wrap *gomock.Call +type MockPacketConnSetDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnSetDeadlineCall) Return(arg0 error) *MockPacketConnSetDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnSetDeadlineCall) Do(f func(time.Time) error) *MockPacketConnSetDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnSetDeadlineCall) DoAndReturn(f func(time.Time) error) *MockPacketConnSetDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetReadDeadline mocks base method. +func (m *MockPacketConn) SetReadDeadline(t time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetReadDeadline", t) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetReadDeadline indicates an expected call of SetReadDeadline. +func (mr *MockPacketConnMockRecorder) SetReadDeadline(t any) *MockPacketConnSetReadDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetReadDeadline", reflect.TypeOf((*MockPacketConn)(nil).SetReadDeadline), t) + return &MockPacketConnSetReadDeadlineCall{Call: call} +} + +// MockPacketConnSetReadDeadlineCall wrap *gomock.Call +type MockPacketConnSetReadDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnSetReadDeadlineCall) Return(arg0 error) *MockPacketConnSetReadDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnSetReadDeadlineCall) Do(f func(time.Time) error) *MockPacketConnSetReadDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnSetReadDeadlineCall) DoAndReturn(f func(time.Time) error) *MockPacketConnSetReadDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetWriteDeadline mocks base method. +func (m *MockPacketConn) SetWriteDeadline(t time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetWriteDeadline", t) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetWriteDeadline indicates an expected call of SetWriteDeadline. +func (mr *MockPacketConnMockRecorder) SetWriteDeadline(t any) *MockPacketConnSetWriteDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetWriteDeadline", reflect.TypeOf((*MockPacketConn)(nil).SetWriteDeadline), t) + return &MockPacketConnSetWriteDeadlineCall{Call: call} +} + +// MockPacketConnSetWriteDeadlineCall wrap *gomock.Call +type MockPacketConnSetWriteDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnSetWriteDeadlineCall) Return(arg0 error) *MockPacketConnSetWriteDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnSetWriteDeadlineCall) Do(f func(time.Time) error) *MockPacketConnSetWriteDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnSetWriteDeadlineCall) DoAndReturn(f func(time.Time) error) *MockPacketConnSetWriteDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// WriteTo mocks base method. +func (m *MockPacketConn) WriteTo(p []byte, addr net.Addr) (int, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "WriteTo", p, addr) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// WriteTo indicates an expected call of WriteTo. +func (mr *MockPacketConnMockRecorder) WriteTo(p, addr any) *MockPacketConnWriteToCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WriteTo", reflect.TypeOf((*MockPacketConn)(nil).WriteTo), p, addr) + return &MockPacketConnWriteToCall{Call: call} +} + +// MockPacketConnWriteToCall wrap *gomock.Call +type MockPacketConnWriteToCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockPacketConnWriteToCall) Return(n int, err error) *MockPacketConnWriteToCall { + c.Call = c.Call.Return(n, err) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockPacketConnWriteToCall) Do(f func([]byte, net.Addr) (int, error)) *MockPacketConnWriteToCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockPacketConnWriteToCall) DoAndReturn(f func([]byte, net.Addr) (int, error)) *MockPacketConnWriteToCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_raw_conn_test.go b/third_party/quic-go/mock_raw_conn_test.go new file mode 100644 index 0000000..4ac1697 --- /dev/null +++ b/third_party/quic-go/mock_raw_conn_test.go @@ -0,0 +1,273 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: RawConn) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_raw_conn_test.go github.com/apernet/quic-go RawConn +// + +// Package quic is a generated GoMock package. +package quic + +import ( + net "net" + reflect "reflect" + time "time" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockRawConn is a mock of RawConn interface. +type MockRawConn struct { + ctrl *gomock.Controller + recorder *MockRawConnMockRecorder + isgomock struct{} +} + +// MockRawConnMockRecorder is the mock recorder for MockRawConn. +type MockRawConnMockRecorder struct { + mock *MockRawConn +} + +// NewMockRawConn creates a new mock instance. +func NewMockRawConn(ctrl *gomock.Controller) *MockRawConn { + mock := &MockRawConn{ctrl: ctrl} + mock.recorder = &MockRawConnMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockRawConn) EXPECT() *MockRawConnMockRecorder { + return m.recorder +} + +// Close mocks base method. +func (m *MockRawConn) Close() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Close") + ret0, _ := ret[0].(error) + return ret0 +} + +// Close indicates an expected call of Close. +func (mr *MockRawConnMockRecorder) Close() *MockRawConnCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockRawConn)(nil).Close)) + return &MockRawConnCloseCall{Call: call} +} + +// MockRawConnCloseCall wrap *gomock.Call +type MockRawConnCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockRawConnCloseCall) Return(arg0 error) *MockRawConnCloseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockRawConnCloseCall) Do(f func() error) *MockRawConnCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockRawConnCloseCall) DoAndReturn(f func() error) *MockRawConnCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// LocalAddr mocks base method. +func (m *MockRawConn) LocalAddr() net.Addr { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "LocalAddr") + ret0, _ := ret[0].(net.Addr) + return ret0 +} + +// LocalAddr indicates an expected call of LocalAddr. +func (mr *MockRawConnMockRecorder) LocalAddr() *MockRawConnLocalAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LocalAddr", reflect.TypeOf((*MockRawConn)(nil).LocalAddr)) + return &MockRawConnLocalAddrCall{Call: call} +} + +// MockRawConnLocalAddrCall wrap *gomock.Call +type MockRawConnLocalAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockRawConnLocalAddrCall) Return(arg0 net.Addr) *MockRawConnLocalAddrCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockRawConnLocalAddrCall) Do(f func() net.Addr) *MockRawConnLocalAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockRawConnLocalAddrCall) DoAndReturn(f func() net.Addr) *MockRawConnLocalAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// ReadPacket mocks base method. +func (m *MockRawConn) ReadPacket() (receivedPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReadPacket") + ret0, _ := ret[0].(receivedPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ReadPacket indicates an expected call of ReadPacket. +func (mr *MockRawConnMockRecorder) ReadPacket() *MockRawConnReadPacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReadPacket", reflect.TypeOf((*MockRawConn)(nil).ReadPacket)) + return &MockRawConnReadPacketCall{Call: call} +} + +// MockRawConnReadPacketCall wrap *gomock.Call +type MockRawConnReadPacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockRawConnReadPacketCall) Return(arg0 receivedPacket, arg1 error) *MockRawConnReadPacketCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockRawConnReadPacketCall) Do(f func() (receivedPacket, error)) *MockRawConnReadPacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockRawConnReadPacketCall) DoAndReturn(f func() (receivedPacket, error)) *MockRawConnReadPacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetReadDeadline mocks base method. +func (m *MockRawConn) SetReadDeadline(arg0 time.Time) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetReadDeadline", arg0) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetReadDeadline indicates an expected call of SetReadDeadline. +func (mr *MockRawConnMockRecorder) SetReadDeadline(arg0 any) *MockRawConnSetReadDeadlineCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetReadDeadline", reflect.TypeOf((*MockRawConn)(nil).SetReadDeadline), arg0) + return &MockRawConnSetReadDeadlineCall{Call: call} +} + +// MockRawConnSetReadDeadlineCall wrap *gomock.Call +type MockRawConnSetReadDeadlineCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockRawConnSetReadDeadlineCall) Return(arg0 error) *MockRawConnSetReadDeadlineCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockRawConnSetReadDeadlineCall) Do(f func(time.Time) error) *MockRawConnSetReadDeadlineCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockRawConnSetReadDeadlineCall) DoAndReturn(f func(time.Time) error) *MockRawConnSetReadDeadlineCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// WritePacket mocks base method. +func (m *MockRawConn) WritePacket(b []byte, addr net.Addr, packetInfoOOB []byte, gsoSize uint16, ecn protocol.ECN) (int, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "WritePacket", b, addr, packetInfoOOB, gsoSize, ecn) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// WritePacket indicates an expected call of WritePacket. +func (mr *MockRawConnMockRecorder) WritePacket(b, addr, packetInfoOOB, gsoSize, ecn any) *MockRawConnWritePacketCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WritePacket", reflect.TypeOf((*MockRawConn)(nil).WritePacket), b, addr, packetInfoOOB, gsoSize, ecn) + return &MockRawConnWritePacketCall{Call: call} +} + +// MockRawConnWritePacketCall wrap *gomock.Call +type MockRawConnWritePacketCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockRawConnWritePacketCall) Return(arg0 int, arg1 error) *MockRawConnWritePacketCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockRawConnWritePacketCall) Do(f func([]byte, net.Addr, []byte, uint16, protocol.ECN) (int, error)) *MockRawConnWritePacketCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockRawConnWritePacketCall) DoAndReturn(f func([]byte, net.Addr, []byte, uint16, protocol.ECN) (int, error)) *MockRawConnWritePacketCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// capabilities mocks base method. +func (m *MockRawConn) capabilities() connCapabilities { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "capabilities") + ret0, _ := ret[0].(connCapabilities) + return ret0 +} + +// capabilities indicates an expected call of capabilities. +func (mr *MockRawConnMockRecorder) capabilities() *MockRawConncapabilitiesCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "capabilities", reflect.TypeOf((*MockRawConn)(nil).capabilities)) + return &MockRawConncapabilitiesCall{Call: call} +} + +// MockRawConncapabilitiesCall wrap *gomock.Call +type MockRawConncapabilitiesCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockRawConncapabilitiesCall) Return(arg0 connCapabilities) *MockRawConncapabilitiesCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockRawConncapabilitiesCall) Do(f func() connCapabilities) *MockRawConncapabilitiesCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockRawConncapabilitiesCall) DoAndReturn(f func() connCapabilities) *MockRawConncapabilitiesCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_sealing_manager_test.go b/third_party/quic-go/mock_sealing_manager_test.go new file mode 100644 index 0000000..aa2b117 --- /dev/null +++ b/third_party/quic-go/mock_sealing_manager_test.go @@ -0,0 +1,197 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: SealingManager) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_sealing_manager_test.go github.com/apernet/quic-go SealingManager +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + handshake "github.com/apernet/quic-go/internal/handshake" + gomock "go.uber.org/mock/gomock" +) + +// MockSealingManager is a mock of SealingManager interface. +type MockSealingManager struct { + ctrl *gomock.Controller + recorder *MockSealingManagerMockRecorder + isgomock struct{} +} + +// MockSealingManagerMockRecorder is the mock recorder for MockSealingManager. +type MockSealingManagerMockRecorder struct { + mock *MockSealingManager +} + +// NewMockSealingManager creates a new mock instance. +func NewMockSealingManager(ctrl *gomock.Controller) *MockSealingManager { + mock := &MockSealingManager{ctrl: ctrl} + mock.recorder = &MockSealingManagerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockSealingManager) EXPECT() *MockSealingManagerMockRecorder { + return m.recorder +} + +// Get0RTTSealer mocks base method. +func (m *MockSealingManager) Get0RTTSealer() (handshake.LongHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Get0RTTSealer") + ret0, _ := ret[0].(handshake.LongHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Get0RTTSealer indicates an expected call of Get0RTTSealer. +func (mr *MockSealingManagerMockRecorder) Get0RTTSealer() *MockSealingManagerGet0RTTSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get0RTTSealer", reflect.TypeOf((*MockSealingManager)(nil).Get0RTTSealer)) + return &MockSealingManagerGet0RTTSealerCall{Call: call} +} + +// MockSealingManagerGet0RTTSealerCall wrap *gomock.Call +type MockSealingManagerGet0RTTSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSealingManagerGet0RTTSealerCall) Return(arg0 handshake.LongHeaderSealer, arg1 error) *MockSealingManagerGet0RTTSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSealingManagerGet0RTTSealerCall) Do(f func() (handshake.LongHeaderSealer, error)) *MockSealingManagerGet0RTTSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSealingManagerGet0RTTSealerCall) DoAndReturn(f func() (handshake.LongHeaderSealer, error)) *MockSealingManagerGet0RTTSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Get1RTTSealer mocks base method. +func (m *MockSealingManager) Get1RTTSealer() (handshake.ShortHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Get1RTTSealer") + ret0, _ := ret[0].(handshake.ShortHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Get1RTTSealer indicates an expected call of Get1RTTSealer. +func (mr *MockSealingManagerMockRecorder) Get1RTTSealer() *MockSealingManagerGet1RTTSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get1RTTSealer", reflect.TypeOf((*MockSealingManager)(nil).Get1RTTSealer)) + return &MockSealingManagerGet1RTTSealerCall{Call: call} +} + +// MockSealingManagerGet1RTTSealerCall wrap *gomock.Call +type MockSealingManagerGet1RTTSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSealingManagerGet1RTTSealerCall) Return(arg0 handshake.ShortHeaderSealer, arg1 error) *MockSealingManagerGet1RTTSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSealingManagerGet1RTTSealerCall) Do(f func() (handshake.ShortHeaderSealer, error)) *MockSealingManagerGet1RTTSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSealingManagerGet1RTTSealerCall) DoAndReturn(f func() (handshake.ShortHeaderSealer, error)) *MockSealingManagerGet1RTTSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetHandshakeSealer mocks base method. +func (m *MockSealingManager) GetHandshakeSealer() (handshake.LongHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetHandshakeSealer") + ret0, _ := ret[0].(handshake.LongHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetHandshakeSealer indicates an expected call of GetHandshakeSealer. +func (mr *MockSealingManagerMockRecorder) GetHandshakeSealer() *MockSealingManagerGetHandshakeSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetHandshakeSealer", reflect.TypeOf((*MockSealingManager)(nil).GetHandshakeSealer)) + return &MockSealingManagerGetHandshakeSealerCall{Call: call} +} + +// MockSealingManagerGetHandshakeSealerCall wrap *gomock.Call +type MockSealingManagerGetHandshakeSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSealingManagerGetHandshakeSealerCall) Return(arg0 handshake.LongHeaderSealer, arg1 error) *MockSealingManagerGetHandshakeSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSealingManagerGetHandshakeSealerCall) Do(f func() (handshake.LongHeaderSealer, error)) *MockSealingManagerGetHandshakeSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSealingManagerGetHandshakeSealerCall) DoAndReturn(f func() (handshake.LongHeaderSealer, error)) *MockSealingManagerGetHandshakeSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// GetInitialSealer mocks base method. +func (m *MockSealingManager) GetInitialSealer() (handshake.LongHeaderSealer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetInitialSealer") + ret0, _ := ret[0].(handshake.LongHeaderSealer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetInitialSealer indicates an expected call of GetInitialSealer. +func (mr *MockSealingManagerMockRecorder) GetInitialSealer() *MockSealingManagerGetInitialSealerCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetInitialSealer", reflect.TypeOf((*MockSealingManager)(nil).GetInitialSealer)) + return &MockSealingManagerGetInitialSealerCall{Call: call} +} + +// MockSealingManagerGetInitialSealerCall wrap *gomock.Call +type MockSealingManagerGetInitialSealerCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSealingManagerGetInitialSealerCall) Return(arg0 handshake.LongHeaderSealer, arg1 error) *MockSealingManagerGetInitialSealerCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSealingManagerGetInitialSealerCall) Do(f func() (handshake.LongHeaderSealer, error)) *MockSealingManagerGetInitialSealerCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSealingManagerGetInitialSealerCall) DoAndReturn(f func() (handshake.LongHeaderSealer, error)) *MockSealingManagerGetInitialSealerCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_send_conn_test.go b/third_party/quic-go/mock_send_conn_test.go new file mode 100644 index 0000000..2bc871e --- /dev/null +++ b/third_party/quic-go/mock_send_conn_test.go @@ -0,0 +1,342 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: SendConn) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_send_conn_test.go github.com/apernet/quic-go SendConn +// + +// Package quic is a generated GoMock package. +package quic + +import ( + net "net" + reflect "reflect" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockSendConn is a mock of SendConn interface. +type MockSendConn struct { + ctrl *gomock.Controller + recorder *MockSendConnMockRecorder + isgomock struct{} +} + +// MockSendConnMockRecorder is the mock recorder for MockSendConn. +type MockSendConnMockRecorder struct { + mock *MockSendConn +} + +// NewMockSendConn creates a new mock instance. +func NewMockSendConn(ctrl *gomock.Controller) *MockSendConn { + mock := &MockSendConn{ctrl: ctrl} + mock.recorder = &MockSendConnMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockSendConn) EXPECT() *MockSendConnMockRecorder { + return m.recorder +} + +// ChangeRemoteAddr mocks base method. +func (m *MockSendConn) ChangeRemoteAddr(addr net.Addr, info packetInfo) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "ChangeRemoteAddr", addr, info) +} + +// ChangeRemoteAddr indicates an expected call of ChangeRemoteAddr. +func (mr *MockSendConnMockRecorder) ChangeRemoteAddr(addr, info any) *MockSendConnChangeRemoteAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ChangeRemoteAddr", reflect.TypeOf((*MockSendConn)(nil).ChangeRemoteAddr), addr, info) + return &MockSendConnChangeRemoteAddrCall{Call: call} +} + +// MockSendConnChangeRemoteAddrCall wrap *gomock.Call +type MockSendConnChangeRemoteAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnChangeRemoteAddrCall) Return() *MockSendConnChangeRemoteAddrCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnChangeRemoteAddrCall) Do(f func(net.Addr, packetInfo)) *MockSendConnChangeRemoteAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnChangeRemoteAddrCall) DoAndReturn(f func(net.Addr, packetInfo)) *MockSendConnChangeRemoteAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Close mocks base method. +func (m *MockSendConn) Close() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Close") + ret0, _ := ret[0].(error) + return ret0 +} + +// Close indicates an expected call of Close. +func (mr *MockSendConnMockRecorder) Close() *MockSendConnCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockSendConn)(nil).Close)) + return &MockSendConnCloseCall{Call: call} +} + +// MockSendConnCloseCall wrap *gomock.Call +type MockSendConnCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnCloseCall) Return(arg0 error) *MockSendConnCloseCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnCloseCall) Do(f func() error) *MockSendConnCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnCloseCall) DoAndReturn(f func() error) *MockSendConnCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// LocalAddr mocks base method. +func (m *MockSendConn) LocalAddr() net.Addr { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "LocalAddr") + ret0, _ := ret[0].(net.Addr) + return ret0 +} + +// LocalAddr indicates an expected call of LocalAddr. +func (mr *MockSendConnMockRecorder) LocalAddr() *MockSendConnLocalAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LocalAddr", reflect.TypeOf((*MockSendConn)(nil).LocalAddr)) + return &MockSendConnLocalAddrCall{Call: call} +} + +// MockSendConnLocalAddrCall wrap *gomock.Call +type MockSendConnLocalAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnLocalAddrCall) Return(arg0 net.Addr) *MockSendConnLocalAddrCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnLocalAddrCall) Do(f func() net.Addr) *MockSendConnLocalAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnLocalAddrCall) DoAndReturn(f func() net.Addr) *MockSendConnLocalAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// RemoteAddr mocks base method. +func (m *MockSendConn) RemoteAddr() net.Addr { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "RemoteAddr") + ret0, _ := ret[0].(net.Addr) + return ret0 +} + +// RemoteAddr indicates an expected call of RemoteAddr. +func (mr *MockSendConnMockRecorder) RemoteAddr() *MockSendConnRemoteAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RemoteAddr", reflect.TypeOf((*MockSendConn)(nil).RemoteAddr)) + return &MockSendConnRemoteAddrCall{Call: call} +} + +// MockSendConnRemoteAddrCall wrap *gomock.Call +type MockSendConnRemoteAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnRemoteAddrCall) Return(arg0 net.Addr) *MockSendConnRemoteAddrCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnRemoteAddrCall) Do(f func() net.Addr) *MockSendConnRemoteAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnRemoteAddrCall) DoAndReturn(f func() net.Addr) *MockSendConnRemoteAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SetRemoteAddr mocks base method. +func (m *MockSendConn) SetRemoteAddr(addr net.Addr) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SetRemoteAddr", addr) +} + +// SetRemoteAddr indicates an expected call of SetRemoteAddr. +func (mr *MockSendConnMockRecorder) SetRemoteAddr(addr any) *MockSendConnSetRemoteAddrCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetRemoteAddr", reflect.TypeOf((*MockSendConn)(nil).SetRemoteAddr), addr) + return &MockSendConnSetRemoteAddrCall{Call: call} +} + +// MockSendConnSetRemoteAddrCall wrap *gomock.Call +type MockSendConnSetRemoteAddrCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnSetRemoteAddrCall) Return() *MockSendConnSetRemoteAddrCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnSetRemoteAddrCall) Do(f func(net.Addr)) *MockSendConnSetRemoteAddrCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnSetRemoteAddrCall) DoAndReturn(f func(net.Addr)) *MockSendConnSetRemoteAddrCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Write mocks base method. +func (m *MockSendConn) Write(b []byte, gsoSize uint16, ecn protocol.ECN) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Write", b, gsoSize, ecn) + ret0, _ := ret[0].(error) + return ret0 +} + +// Write indicates an expected call of Write. +func (mr *MockSendConnMockRecorder) Write(b, gsoSize, ecn any) *MockSendConnWriteCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Write", reflect.TypeOf((*MockSendConn)(nil).Write), b, gsoSize, ecn) + return &MockSendConnWriteCall{Call: call} +} + +// MockSendConnWriteCall wrap *gomock.Call +type MockSendConnWriteCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnWriteCall) Return(arg0 error) *MockSendConnWriteCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnWriteCall) Do(f func([]byte, uint16, protocol.ECN) error) *MockSendConnWriteCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnWriteCall) DoAndReturn(f func([]byte, uint16, protocol.ECN) error) *MockSendConnWriteCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// WriteTo mocks base method. +func (m *MockSendConn) WriteTo(arg0 []byte, arg1 net.Addr, arg2 packetInfo) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "WriteTo", arg0, arg1, arg2) + ret0, _ := ret[0].(error) + return ret0 +} + +// WriteTo indicates an expected call of WriteTo. +func (mr *MockSendConnMockRecorder) WriteTo(arg0, arg1, arg2 any) *MockSendConnWriteToCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WriteTo", reflect.TypeOf((*MockSendConn)(nil).WriteTo), arg0, arg1, arg2) + return &MockSendConnWriteToCall{Call: call} +} + +// MockSendConnWriteToCall wrap *gomock.Call +type MockSendConnWriteToCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConnWriteToCall) Return(arg0 error) *MockSendConnWriteToCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConnWriteToCall) Do(f func([]byte, net.Addr, packetInfo) error) *MockSendConnWriteToCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConnWriteToCall) DoAndReturn(f func([]byte, net.Addr, packetInfo) error) *MockSendConnWriteToCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// capabilities mocks base method. +func (m *MockSendConn) capabilities() connCapabilities { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "capabilities") + ret0, _ := ret[0].(connCapabilities) + return ret0 +} + +// capabilities indicates an expected call of capabilities. +func (mr *MockSendConnMockRecorder) capabilities() *MockSendConncapabilitiesCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "capabilities", reflect.TypeOf((*MockSendConn)(nil).capabilities)) + return &MockSendConncapabilitiesCall{Call: call} +} + +// MockSendConncapabilitiesCall wrap *gomock.Call +type MockSendConncapabilitiesCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSendConncapabilitiesCall) Return(arg0 connCapabilities) *MockSendConncapabilitiesCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSendConncapabilitiesCall) Do(f func() connCapabilities) *MockSendConncapabilitiesCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSendConncapabilitiesCall) DoAndReturn(f func() connCapabilities) *MockSendConncapabilitiesCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_sender_test.go b/third_party/quic-go/mock_sender_test.go new file mode 100644 index 0000000..c48e8d7 --- /dev/null +++ b/third_party/quic-go/mock_sender_test.go @@ -0,0 +1,264 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: Sender) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_sender_test.go github.com/apernet/quic-go Sender +// + +// Package quic is a generated GoMock package. +package quic + +import ( + net "net" + reflect "reflect" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockSender is a mock of Sender interface. +type MockSender struct { + ctrl *gomock.Controller + recorder *MockSenderMockRecorder + isgomock struct{} +} + +// MockSenderMockRecorder is the mock recorder for MockSender. +type MockSenderMockRecorder struct { + mock *MockSender +} + +// NewMockSender creates a new mock instance. +func NewMockSender(ctrl *gomock.Controller) *MockSender { + mock := &MockSender{ctrl: ctrl} + mock.recorder = &MockSenderMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockSender) EXPECT() *MockSenderMockRecorder { + return m.recorder +} + +// Available mocks base method. +func (m *MockSender) Available() <-chan struct{} { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Available") + ret0, _ := ret[0].(<-chan struct{}) + return ret0 +} + +// Available indicates an expected call of Available. +func (mr *MockSenderMockRecorder) Available() *MockSenderAvailableCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Available", reflect.TypeOf((*MockSender)(nil).Available)) + return &MockSenderAvailableCall{Call: call} +} + +// MockSenderAvailableCall wrap *gomock.Call +type MockSenderAvailableCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSenderAvailableCall) Return(arg0 <-chan struct{}) *MockSenderAvailableCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSenderAvailableCall) Do(f func() <-chan struct{}) *MockSenderAvailableCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSenderAvailableCall) DoAndReturn(f func() <-chan struct{}) *MockSenderAvailableCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Close mocks base method. +func (m *MockSender) Close() { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Close") +} + +// Close indicates an expected call of Close. +func (mr *MockSenderMockRecorder) Close() *MockSenderCloseCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockSender)(nil).Close)) + return &MockSenderCloseCall{Call: call} +} + +// MockSenderCloseCall wrap *gomock.Call +type MockSenderCloseCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSenderCloseCall) Return() *MockSenderCloseCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSenderCloseCall) Do(f func()) *MockSenderCloseCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSenderCloseCall) DoAndReturn(f func()) *MockSenderCloseCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Run mocks base method. +func (m *MockSender) Run() error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Run") + ret0, _ := ret[0].(error) + return ret0 +} + +// Run indicates an expected call of Run. +func (mr *MockSenderMockRecorder) Run() *MockSenderRunCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Run", reflect.TypeOf((*MockSender)(nil).Run)) + return &MockSenderRunCall{Call: call} +} + +// MockSenderRunCall wrap *gomock.Call +type MockSenderRunCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSenderRunCall) Return(arg0 error) *MockSenderRunCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSenderRunCall) Do(f func() error) *MockSenderRunCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSenderRunCall) DoAndReturn(f func() error) *MockSenderRunCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// Send mocks base method. +func (m *MockSender) Send(p *packetBuffer, gsoSize uint16, ecn protocol.ECN) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Send", p, gsoSize, ecn) +} + +// Send indicates an expected call of Send. +func (mr *MockSenderMockRecorder) Send(p, gsoSize, ecn any) *MockSenderSendCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Send", reflect.TypeOf((*MockSender)(nil).Send), p, gsoSize, ecn) + return &MockSenderSendCall{Call: call} +} + +// MockSenderSendCall wrap *gomock.Call +type MockSenderSendCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSenderSendCall) Return() *MockSenderSendCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSenderSendCall) Do(f func(*packetBuffer, uint16, protocol.ECN)) *MockSenderSendCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSenderSendCall) DoAndReturn(f func(*packetBuffer, uint16, protocol.ECN)) *MockSenderSendCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// SendProbe mocks base method. +func (m *MockSender) SendProbe(arg0 *packetBuffer, arg1 net.Addr, arg2 packetInfo) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "SendProbe", arg0, arg1, arg2) +} + +// SendProbe indicates an expected call of SendProbe. +func (mr *MockSenderMockRecorder) SendProbe(arg0, arg1, arg2 any) *MockSenderSendProbeCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SendProbe", reflect.TypeOf((*MockSender)(nil).SendProbe), arg0, arg1, arg2) + return &MockSenderSendProbeCall{Call: call} +} + +// MockSenderSendProbeCall wrap *gomock.Call +type MockSenderSendProbeCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSenderSendProbeCall) Return() *MockSenderSendProbeCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSenderSendProbeCall) Do(f func(*packetBuffer, net.Addr, packetInfo)) *MockSenderSendProbeCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSenderSendProbeCall) DoAndReturn(f func(*packetBuffer, net.Addr, packetInfo)) *MockSenderSendProbeCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// WouldBlock mocks base method. +func (m *MockSender) WouldBlock() bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "WouldBlock") + ret0, _ := ret[0].(bool) + return ret0 +} + +// WouldBlock indicates an expected call of WouldBlock. +func (mr *MockSenderMockRecorder) WouldBlock() *MockSenderWouldBlockCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WouldBlock", reflect.TypeOf((*MockSender)(nil).WouldBlock)) + return &MockSenderWouldBlockCall{Call: call} +} + +// MockSenderWouldBlockCall wrap *gomock.Call +type MockSenderWouldBlockCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockSenderWouldBlockCall) Return(arg0 bool) *MockSenderWouldBlockCall { + c.Call = c.Call.Return(arg0) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockSenderWouldBlockCall) Do(f func() bool) *MockSenderWouldBlockCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockSenderWouldBlockCall) DoAndReturn(f func() bool) *MockSenderWouldBlockCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_stream_control_frame_getter_test.go b/third_party/quic-go/mock_stream_control_frame_getter_test.go new file mode 100644 index 0000000..8155aac --- /dev/null +++ b/third_party/quic-go/mock_stream_control_frame_getter_test.go @@ -0,0 +1,82 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: StreamControlFrameGetter) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_stream_control_frame_getter_test.go github.com/apernet/quic-go StreamControlFrameGetter +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + ackhandler "github.com/apernet/quic-go/internal/ackhandler" + monotime "github.com/apernet/quic-go/internal/monotime" + gomock "go.uber.org/mock/gomock" +) + +// MockStreamControlFrameGetter is a mock of StreamControlFrameGetter interface. +type MockStreamControlFrameGetter struct { + ctrl *gomock.Controller + recorder *MockStreamControlFrameGetterMockRecorder + isgomock struct{} +} + +// MockStreamControlFrameGetterMockRecorder is the mock recorder for MockStreamControlFrameGetter. +type MockStreamControlFrameGetterMockRecorder struct { + mock *MockStreamControlFrameGetter +} + +// NewMockStreamControlFrameGetter creates a new mock instance. +func NewMockStreamControlFrameGetter(ctrl *gomock.Controller) *MockStreamControlFrameGetter { + mock := &MockStreamControlFrameGetter{ctrl: ctrl} + mock.recorder = &MockStreamControlFrameGetterMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockStreamControlFrameGetter) EXPECT() *MockStreamControlFrameGetterMockRecorder { + return m.recorder +} + +// getControlFrame mocks base method. +func (m *MockStreamControlFrameGetter) getControlFrame(arg0 monotime.Time) (ackhandler.Frame, bool, bool) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "getControlFrame", arg0) + ret0, _ := ret[0].(ackhandler.Frame) + ret1, _ := ret[1].(bool) + ret2, _ := ret[2].(bool) + return ret0, ret1, ret2 +} + +// getControlFrame indicates an expected call of getControlFrame. +func (mr *MockStreamControlFrameGetterMockRecorder) getControlFrame(arg0 any) *MockStreamControlFrameGettergetControlFrameCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "getControlFrame", reflect.TypeOf((*MockStreamControlFrameGetter)(nil).getControlFrame), arg0) + return &MockStreamControlFrameGettergetControlFrameCall{Call: call} +} + +// MockStreamControlFrameGettergetControlFrameCall wrap *gomock.Call +type MockStreamControlFrameGettergetControlFrameCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockStreamControlFrameGettergetControlFrameCall) Return(arg0 ackhandler.Frame, ok, hasMore bool) *MockStreamControlFrameGettergetControlFrameCall { + c.Call = c.Call.Return(arg0, ok, hasMore) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockStreamControlFrameGettergetControlFrameCall) Do(f func(monotime.Time) (ackhandler.Frame, bool, bool)) *MockStreamControlFrameGettergetControlFrameCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockStreamControlFrameGettergetControlFrameCall) DoAndReturn(f func(monotime.Time) (ackhandler.Frame, bool, bool)) *MockStreamControlFrameGettergetControlFrameCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_stream_frame_getter_test.go b/third_party/quic-go/mock_stream_frame_getter_test.go new file mode 100644 index 0000000..809acb8 --- /dev/null +++ b/third_party/quic-go/mock_stream_frame_getter_test.go @@ -0,0 +1,83 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: StreamFrameGetter) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_stream_frame_getter_test.go github.com/apernet/quic-go StreamFrameGetter +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + ackhandler "github.com/apernet/quic-go/internal/ackhandler" + protocol "github.com/apernet/quic-go/internal/protocol" + wire "github.com/apernet/quic-go/internal/wire" + gomock "go.uber.org/mock/gomock" +) + +// MockStreamFrameGetter is a mock of StreamFrameGetter interface. +type MockStreamFrameGetter struct { + ctrl *gomock.Controller + recorder *MockStreamFrameGetterMockRecorder + isgomock struct{} +} + +// MockStreamFrameGetterMockRecorder is the mock recorder for MockStreamFrameGetter. +type MockStreamFrameGetterMockRecorder struct { + mock *MockStreamFrameGetter +} + +// NewMockStreamFrameGetter creates a new mock instance. +func NewMockStreamFrameGetter(ctrl *gomock.Controller) *MockStreamFrameGetter { + mock := &MockStreamFrameGetter{ctrl: ctrl} + mock.recorder = &MockStreamFrameGetterMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockStreamFrameGetter) EXPECT() *MockStreamFrameGetterMockRecorder { + return m.recorder +} + +// popStreamFrame mocks base method. +func (m *MockStreamFrameGetter) popStreamFrame(arg0 protocol.ByteCount, arg1 protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "popStreamFrame", arg0, arg1) + ret0, _ := ret[0].(ackhandler.StreamFrame) + ret1, _ := ret[1].(*wire.StreamDataBlockedFrame) + ret2, _ := ret[2].(bool) + return ret0, ret1, ret2 +} + +// popStreamFrame indicates an expected call of popStreamFrame. +func (mr *MockStreamFrameGetterMockRecorder) popStreamFrame(arg0, arg1 any) *MockStreamFrameGetterpopStreamFrameCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "popStreamFrame", reflect.TypeOf((*MockStreamFrameGetter)(nil).popStreamFrame), arg0, arg1) + return &MockStreamFrameGetterpopStreamFrameCall{Call: call} +} + +// MockStreamFrameGetterpopStreamFrameCall wrap *gomock.Call +type MockStreamFrameGetterpopStreamFrameCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockStreamFrameGetterpopStreamFrameCall) Return(arg0 ackhandler.StreamFrame, arg1 *wire.StreamDataBlockedFrame, arg2 bool) *MockStreamFrameGetterpopStreamFrameCall { + c.Call = c.Call.Return(arg0, arg1, arg2) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockStreamFrameGetterpopStreamFrameCall) Do(f func(protocol.ByteCount, protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool)) *MockStreamFrameGetterpopStreamFrameCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockStreamFrameGetterpopStreamFrameCall) DoAndReturn(f func(protocol.ByteCount, protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool)) *MockStreamFrameGetterpopStreamFrameCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_stream_sender_test.go b/third_party/quic-go/mock_stream_sender_test.go new file mode 100644 index 0000000..8758d84 --- /dev/null +++ b/third_party/quic-go/mock_stream_sender_test.go @@ -0,0 +1,185 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: StreamSender) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_stream_sender_test.go github.com/apernet/quic-go StreamSender +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + protocol "github.com/apernet/quic-go/internal/protocol" + gomock "go.uber.org/mock/gomock" +) + +// MockStreamSender is a mock of StreamSender interface. +type MockStreamSender struct { + ctrl *gomock.Controller + recorder *MockStreamSenderMockRecorder + isgomock struct{} +} + +// MockStreamSenderMockRecorder is the mock recorder for MockStreamSender. +type MockStreamSenderMockRecorder struct { + mock *MockStreamSender +} + +// NewMockStreamSender creates a new mock instance. +func NewMockStreamSender(ctrl *gomock.Controller) *MockStreamSender { + mock := &MockStreamSender{ctrl: ctrl} + mock.recorder = &MockStreamSenderMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockStreamSender) EXPECT() *MockStreamSenderMockRecorder { + return m.recorder +} + +// onHasConnectionData mocks base method. +func (m *MockStreamSender) onHasConnectionData() { + m.ctrl.T.Helper() + m.ctrl.Call(m, "onHasConnectionData") +} + +// onHasConnectionData indicates an expected call of onHasConnectionData. +func (mr *MockStreamSenderMockRecorder) onHasConnectionData() *MockStreamSenderonHasConnectionDataCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "onHasConnectionData", reflect.TypeOf((*MockStreamSender)(nil).onHasConnectionData)) + return &MockStreamSenderonHasConnectionDataCall{Call: call} +} + +// MockStreamSenderonHasConnectionDataCall wrap *gomock.Call +type MockStreamSenderonHasConnectionDataCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockStreamSenderonHasConnectionDataCall) Return() *MockStreamSenderonHasConnectionDataCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockStreamSenderonHasConnectionDataCall) Do(f func()) *MockStreamSenderonHasConnectionDataCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockStreamSenderonHasConnectionDataCall) DoAndReturn(f func()) *MockStreamSenderonHasConnectionDataCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// onHasStreamControlFrame mocks base method. +func (m *MockStreamSender) onHasStreamControlFrame(arg0 protocol.StreamID, arg1 streamControlFrameGetter) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "onHasStreamControlFrame", arg0, arg1) +} + +// onHasStreamControlFrame indicates an expected call of onHasStreamControlFrame. +func (mr *MockStreamSenderMockRecorder) onHasStreamControlFrame(arg0, arg1 any) *MockStreamSenderonHasStreamControlFrameCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "onHasStreamControlFrame", reflect.TypeOf((*MockStreamSender)(nil).onHasStreamControlFrame), arg0, arg1) + return &MockStreamSenderonHasStreamControlFrameCall{Call: call} +} + +// MockStreamSenderonHasStreamControlFrameCall wrap *gomock.Call +type MockStreamSenderonHasStreamControlFrameCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockStreamSenderonHasStreamControlFrameCall) Return() *MockStreamSenderonHasStreamControlFrameCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockStreamSenderonHasStreamControlFrameCall) Do(f func(protocol.StreamID, streamControlFrameGetter)) *MockStreamSenderonHasStreamControlFrameCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockStreamSenderonHasStreamControlFrameCall) DoAndReturn(f func(protocol.StreamID, streamControlFrameGetter)) *MockStreamSenderonHasStreamControlFrameCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// onHasStreamData mocks base method. +func (m *MockStreamSender) onHasStreamData(arg0 protocol.StreamID, arg1 *SendStream) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "onHasStreamData", arg0, arg1) +} + +// onHasStreamData indicates an expected call of onHasStreamData. +func (mr *MockStreamSenderMockRecorder) onHasStreamData(arg0, arg1 any) *MockStreamSenderonHasStreamDataCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "onHasStreamData", reflect.TypeOf((*MockStreamSender)(nil).onHasStreamData), arg0, arg1) + return &MockStreamSenderonHasStreamDataCall{Call: call} +} + +// MockStreamSenderonHasStreamDataCall wrap *gomock.Call +type MockStreamSenderonHasStreamDataCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockStreamSenderonHasStreamDataCall) Return() *MockStreamSenderonHasStreamDataCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockStreamSenderonHasStreamDataCall) Do(f func(protocol.StreamID, *SendStream)) *MockStreamSenderonHasStreamDataCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockStreamSenderonHasStreamDataCall) DoAndReturn(f func(protocol.StreamID, *SendStream)) *MockStreamSenderonHasStreamDataCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// onStreamCompleted mocks base method. +func (m *MockStreamSender) onStreamCompleted(arg0 protocol.StreamID) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "onStreamCompleted", arg0) +} + +// onStreamCompleted indicates an expected call of onStreamCompleted. +func (mr *MockStreamSenderMockRecorder) onStreamCompleted(arg0 any) *MockStreamSenderonStreamCompletedCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "onStreamCompleted", reflect.TypeOf((*MockStreamSender)(nil).onStreamCompleted), arg0) + return &MockStreamSenderonStreamCompletedCall{Call: call} +} + +// MockStreamSenderonStreamCompletedCall wrap *gomock.Call +type MockStreamSenderonStreamCompletedCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockStreamSenderonStreamCompletedCall) Return() *MockStreamSenderonStreamCompletedCall { + c.Call = c.Call.Return() + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockStreamSenderonStreamCompletedCall) Do(f func(protocol.StreamID)) *MockStreamSenderonStreamCompletedCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockStreamSenderonStreamCompletedCall) DoAndReturn(f func(protocol.StreamID)) *MockStreamSenderonStreamCompletedCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mock_unpacker_test.go b/third_party/quic-go/mock_unpacker_test.go new file mode 100644 index 0000000..ce98d83 --- /dev/null +++ b/third_party/quic-go/mock_unpacker_test.go @@ -0,0 +1,124 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: github.com/apernet/quic-go (interfaces: Unpacker) +// +// Generated by this command: +// +// mockgen -typed -build_flags=-tags=gomock -package quic -self_package github.com/apernet/quic-go -destination mock_unpacker_test.go github.com/apernet/quic-go Unpacker +// + +// Package quic is a generated GoMock package. +package quic + +import ( + reflect "reflect" + + monotime "github.com/apernet/quic-go/internal/monotime" + protocol "github.com/apernet/quic-go/internal/protocol" + wire "github.com/apernet/quic-go/internal/wire" + gomock "go.uber.org/mock/gomock" +) + +// MockUnpacker is a mock of Unpacker interface. +type MockUnpacker struct { + ctrl *gomock.Controller + recorder *MockUnpackerMockRecorder + isgomock struct{} +} + +// MockUnpackerMockRecorder is the mock recorder for MockUnpacker. +type MockUnpackerMockRecorder struct { + mock *MockUnpacker +} + +// NewMockUnpacker creates a new mock instance. +func NewMockUnpacker(ctrl *gomock.Controller) *MockUnpacker { + mock := &MockUnpacker{ctrl: ctrl} + mock.recorder = &MockUnpackerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockUnpacker) EXPECT() *MockUnpackerMockRecorder { + return m.recorder +} + +// UnpackLongHeader mocks base method. +func (m *MockUnpacker) UnpackLongHeader(hdr *wire.Header, data []byte) (*unpackedPacket, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UnpackLongHeader", hdr, data) + ret0, _ := ret[0].(*unpackedPacket) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UnpackLongHeader indicates an expected call of UnpackLongHeader. +func (mr *MockUnpackerMockRecorder) UnpackLongHeader(hdr, data any) *MockUnpackerUnpackLongHeaderCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnpackLongHeader", reflect.TypeOf((*MockUnpacker)(nil).UnpackLongHeader), hdr, data) + return &MockUnpackerUnpackLongHeaderCall{Call: call} +} + +// MockUnpackerUnpackLongHeaderCall wrap *gomock.Call +type MockUnpackerUnpackLongHeaderCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockUnpackerUnpackLongHeaderCall) Return(arg0 *unpackedPacket, arg1 error) *MockUnpackerUnpackLongHeaderCall { + c.Call = c.Call.Return(arg0, arg1) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockUnpackerUnpackLongHeaderCall) Do(f func(*wire.Header, []byte) (*unpackedPacket, error)) *MockUnpackerUnpackLongHeaderCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockUnpackerUnpackLongHeaderCall) DoAndReturn(f func(*wire.Header, []byte) (*unpackedPacket, error)) *MockUnpackerUnpackLongHeaderCall { + c.Call = c.Call.DoAndReturn(f) + return c +} + +// UnpackShortHeader mocks base method. +func (m *MockUnpacker) UnpackShortHeader(rcvTime monotime.Time, data []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UnpackShortHeader", rcvTime, data) + ret0, _ := ret[0].(protocol.PacketNumber) + ret1, _ := ret[1].(protocol.PacketNumberLen) + ret2, _ := ret[2].(protocol.KeyPhaseBit) + ret3, _ := ret[3].([]byte) + ret4, _ := ret[4].(error) + return ret0, ret1, ret2, ret3, ret4 +} + +// UnpackShortHeader indicates an expected call of UnpackShortHeader. +func (mr *MockUnpackerMockRecorder) UnpackShortHeader(rcvTime, data any) *MockUnpackerUnpackShortHeaderCall { + mr.mock.ctrl.T.Helper() + call := mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnpackShortHeader", reflect.TypeOf((*MockUnpacker)(nil).UnpackShortHeader), rcvTime, data) + return &MockUnpackerUnpackShortHeaderCall{Call: call} +} + +// MockUnpackerUnpackShortHeaderCall wrap *gomock.Call +type MockUnpackerUnpackShortHeaderCall struct { + *gomock.Call +} + +// Return rewrite *gomock.Call.Return +func (c *MockUnpackerUnpackShortHeaderCall) Return(arg0 protocol.PacketNumber, arg1 protocol.PacketNumberLen, arg2 protocol.KeyPhaseBit, arg3 []byte, arg4 error) *MockUnpackerUnpackShortHeaderCall { + c.Call = c.Call.Return(arg0, arg1, arg2, arg3, arg4) + return c +} + +// Do rewrite *gomock.Call.Do +func (c *MockUnpackerUnpackShortHeaderCall) Do(f func(monotime.Time, []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error)) *MockUnpackerUnpackShortHeaderCall { + c.Call = c.Call.Do(f) + return c +} + +// DoAndReturn rewrite *gomock.Call.DoAndReturn +func (c *MockUnpackerUnpackShortHeaderCall) DoAndReturn(f func(monotime.Time, []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error)) *MockUnpackerUnpackShortHeaderCall { + c.Call = c.Call.DoAndReturn(f) + return c +} diff --git a/third_party/quic-go/mockgen.go b/third_party/quic-go/mockgen.go new file mode 100644 index 0000000..acd1673 --- /dev/null +++ b/third_party/quic-go/mockgen.go @@ -0,0 +1,47 @@ +//go:build gomock || generate + +package quic + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_send_conn_test.go github.com/apernet/quic-go SendConn" +type SendConn = sendConn + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_raw_conn_test.go github.com/apernet/quic-go RawConn" +type RawConn = rawConn + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_sender_test.go github.com/apernet/quic-go Sender" +type Sender = sender + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_stream_sender_test.go github.com/apernet/quic-go StreamSender" +type StreamSender = streamSender + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_stream_control_frame_getter_test.go github.com/apernet/quic-go StreamControlFrameGetter" +type StreamControlFrameGetter = streamControlFrameGetter + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_stream_frame_getter_test.go github.com/apernet/quic-go StreamFrameGetter" +type StreamFrameGetter = streamFrameGetter + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_frame_source_test.go github.com/apernet/quic-go FrameSource" +type FrameSource = frameSource + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_ack_frame_source_test.go github.com/apernet/quic-go AckFrameSource" +type AckFrameSource = ackFrameSource + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_sealing_manager_test.go github.com/apernet/quic-go SealingManager" +type SealingManager = sealingManager + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_unpacker_test.go github.com/apernet/quic-go Unpacker" +type Unpacker = unpacker + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_packer_test.go github.com/apernet/quic-go Packer" +type Packer = packer + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_mtu_discoverer_test.go github.com/apernet/quic-go MTUDiscoverer" +type MTUDiscoverer = mtuDiscoverer + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_conn_runner_test.go github.com/apernet/quic-go ConnRunner" +type ConnRunner = connRunner + +//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -package quic -self_package github.com/apernet/quic-go -destination mock_packet_handler_test.go github.com/apernet/quic-go PacketHandler" +type PacketHandler = packetHandler + +//go:generate sh -c "go tool mockgen -typed -package quic -self_package github.com/apernet/quic-go -self_package github.com/apernet/quic-go -destination mock_packetconn_test.go net PacketConn" diff --git a/third_party/quic-go/module_rename.sh b/third_party/quic-go/module_rename.sh new file mode 100644 index 0000000..af4036a --- /dev/null +++ b/third_party/quic-go/module_rename.sh @@ -0,0 +1,35 @@ +#!/usr/bin/env sh +# Rename the Go module path between upstream quic-go and the apernet fork. +# +# ./module_rename.sh upstream -> fork +# ./module_rename.sh -r fork -> upstream +# +# The module directive of the root go.mod is changed with the canonical +# `go mod edit -module`. Everything else (import paths, docs, scripts and +# nested go.mod require/replace directives) is rewritten textually. +# Hidden directories (.git, .github, .clusterfuzzlite, ...) are skipped. +set -eu + +upstream="github.com/quic-go/quic-go" +fork="github.com/apernet/quic-go" + +from="$upstream" +to="$fork" +if [ "${1:-}" = "-r" ] || [ "${1:-}" = "--reverse" ]; then + from="$fork" + to="$upstream" +fi + +# Escape dots so the path is matched literally by sed. +from_re=$(printf '%s' "$from" | sed 's/\./\\./g') + +# 1. Root module directive: use the canonical tool. +go mod edit -module="$to" + +# 2. References everywhere else: imports, docs, scripts and nested go.mod +# require/replace directives. This script is skipped so it doesn't rewrite +# its own upstream/fork path literals. +find . -type d -name '.?*' -prune -o \ + -type f \( -name '*.go' -o -name '*.md' -o -name '*.sh' -o -name '*.mod' \) \ + ! -name module_rename.sh \ + -exec sed -i "s,${from_re},${to},g" {} + diff --git a/third_party/quic-go/monotime/time.go b/third_party/quic-go/monotime/time.go new file mode 100644 index 0000000..95eaebb --- /dev/null +++ b/third_party/quic-go/monotime/time.go @@ -0,0 +1,37 @@ +package monotime + +import ( + "time" + + "github.com/apernet/quic-go/internal/monotime" +) + +// A Time represents an instant in monotonic time. +// Times can be compared using the comparison operators, but the specific +// value is implementation-dependent and should not be relied upon. +// The zero value of Time doesn't have any specific meaning. +type Time = monotime.Time + +// Now returns the current monotonic time. +func Now() Time { + return monotime.Now() +} + +// Since returns the time elapsed since t. It is shorthand for Now().Sub(t). +func Since(t Time) time.Duration { + return monotime.Since(t) +} + +// Until returns the duration until t. +// It is shorthand for t.Sub(Now()). +// If t is in the past, the returned duration will be negative. +func Until(t Time) time.Duration { + return monotime.Until(t) +} + +// FromTime converts a time.Time to a monotonic Time. +// The conversion is relative to the package's start time and may lose +// precision if the time.Time is far from the start time. +func FromTime(t time.Time) Time { + return monotime.FromTime(t) +} diff --git a/third_party/quic-go/mtu_discoverer.go b/third_party/quic-go/mtu_discoverer.go new file mode 100644 index 0000000..19d7e72 --- /dev/null +++ b/third_party/quic-go/mtu_discoverer.go @@ -0,0 +1,253 @@ +package quic + +import ( + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +type mtuDiscoverer interface { + // Start starts the MTU discovery process. + // It's unnecessary to call ShouldSendProbe before that. + Start(now monotime.Time) + ShouldSendProbe(now monotime.Time) bool + CurrentSize() protocol.ByteCount + GetPing(now monotime.Time) (ping ackhandler.Frame, datagramSize protocol.ByteCount) + Reset(now monotime.Time, start, max protocol.ByteCount) +} + +const ( + // At some point, we have to stop searching for a higher MTU. + // We're happy to send a packet that's 10 bytes smaller than the actual MTU. + maxMTUDiff protocol.ByteCount = 20 + // send a probe packet every mtuProbeDelay RTTs + mtuProbeDelay = 5 + // Once maxLostMTUProbes MTU probe packets larger than a certain size are lost, + // MTU discovery won't probe for larger MTUs than this size. + // The algorithm used here is resilient to packet loss of (maxLostMTUProbes - 1) packets. + maxLostMTUProbes = 3 +) + +// The Path MTU is found by sending a larger packet every now and then. +// If the packet is acknowledged, we conclude that the path supports this larger packet size. +// If the packet is lost, this can mean one of two things: +// 1. The path doesn't support this larger packet size, or +// 2. The packet was lost due to packet loss, independent of its size. +// The algorithm used here is resilient to packet loss of (maxLostMTUProbes - 1) packets. +// For simplicty, the following example use maxLostMTUProbes = 2. +// +// Initialization: +// |------------------------------------------------------------------------------| +// min max +// +// The first MTU probe packet will have size (min+max)/2. +// Assume that this packet is acknowledged. We can now move the min marker, +// and continue the search in the resulting interval. +// +// If 1st probe packet acknowledged: +// |---------------------------------------|--------------------------------------| +// min max +// +// If 1st probe packet lost: +// |---------------------------------------|--------------------------------------| +// min lost[0] max +// +// We can't conclude that the path doesn't support this packet size, since the loss of the probe +// packet could have been unrelated to the packet size. A larger probe packet will be sent later on. +// After a loss, the next probe packet has size (min+lost[0])/2. +// Now assume this probe packet is acknowledged: +// +// 2nd probe packet acknowledged: +// |------------------|--------------------|--------------------------------------| +// min lost[0] max +// +// First of all, we conclude that the path supports at least this MTU. That's progress! +// Second, we probe a bit more aggressively with the next probe packet: +// After an acknowledgement, the next probe packet has size (min+max)/2. +// This means we'll send a packet larger than the first probe packet (which was lost). +// +// If 3rd probe packet acknowledged: +// |-------------------------------------------------|----------------------------| +// min max +// +// We can conclude that the loss of the 1st probe packet was not due to its size, and +// continue searching in a much smaller interval now. +// +// If 3rd probe packet lost: +// |------------------|--------------------|---------|----------------------------| +// min lost[0] max +// +// Since in our example numPTOProbes = 2, and we lost 2 packets smaller than max, we +// conclude that this packet size is not supported on the path, and reduce the maximum +// value of the search interval. +// +// MTU discovery concludes once the interval min and max has been narrowed down to maxMTUDiff. + +type mtuFinder struct { + lastProbeTime monotime.Time + + rttStats *utils.RTTStats + + inFlight protocol.ByteCount // the size of the probe packet currently in flight. InvalidByteCount if none is in flight + min protocol.ByteCount + + // on initialization, we treat the maximum size as the first "lost" packet + lost [maxLostMTUProbes]protocol.ByteCount + lastProbeWasLost bool + + // The generation is used to ignore ACKs / losses for probe packets sent before a reset. + // Resets happen when the connection is migrated to a new path. + // We're therefore not concerned about overflows of this counter. + generation uint8 + + qlogger qlogwriter.Recorder +} + +var _ mtuDiscoverer = &mtuFinder{} + +func newMTUDiscoverer( + rttStats *utils.RTTStats, + start, max protocol.ByteCount, + qlogger qlogwriter.Recorder, +) *mtuFinder { + f := &mtuFinder{ + inFlight: protocol.InvalidByteCount, + rttStats: rttStats, + qlogger: qlogger, + } + f.init(start, max) + return f +} + +func (f *mtuFinder) init(start, max protocol.ByteCount) { + f.min = start + for i := range f.lost { + if i == 0 { + f.lost[i] = max + continue + } + f.lost[i] = protocol.InvalidByteCount + } +} + +func (f *mtuFinder) done() bool { + return f.max()-f.min <= maxMTUDiff+1 +} + +func (f *mtuFinder) max() protocol.ByteCount { + for i, v := range f.lost { + if v == protocol.InvalidByteCount { + return f.lost[i-1] + } + } + return f.lost[len(f.lost)-1] +} + +func (f *mtuFinder) Start(now monotime.Time) { + f.lastProbeTime = now // makes sure the first probe packet is not sent immediately +} + +func (f *mtuFinder) ShouldSendProbe(now monotime.Time) bool { + if f.lastProbeTime.IsZero() { + return false + } + if f.inFlight != protocol.InvalidByteCount || f.done() { + return false + } + return !now.Before(f.lastProbeTime.Add(mtuProbeDelay * f.rttStats.SmoothedRTT())) +} + +func (f *mtuFinder) GetPing(now monotime.Time) (ackhandler.Frame, protocol.ByteCount) { + var size protocol.ByteCount + if f.lastProbeWasLost { + size = (f.min + f.lost[0]) / 2 + } else { + size = (f.min + f.max()) / 2 + } + f.lastProbeTime = now + f.inFlight = size + return ackhandler.Frame{ + Frame: &wire.PingFrame{}, + Handler: &mtuFinderAckHandler{mtuFinder: f, generation: f.generation}, + }, size +} + +func (f *mtuFinder) CurrentSize() protocol.ByteCount { + return f.min +} + +func (f *mtuFinder) Reset(now monotime.Time, start, max protocol.ByteCount) { + f.generation++ + f.lastProbeTime = now + f.lastProbeWasLost = false + f.inFlight = protocol.InvalidByteCount + f.init(start, max) +} + +type mtuFinderAckHandler struct { + *mtuFinder + generation uint8 +} + +var _ ackhandler.FrameHandler = &mtuFinderAckHandler{} + +func (h *mtuFinderAckHandler) OnAcked(wire.Frame) { + if h.generation != h.mtuFinder.generation { + // ACK for probe sent before reset + return + } + size := h.inFlight + if size == protocol.InvalidByteCount { + panic("OnAcked callback called although there's no MTU probe packet in flight") + } + h.inFlight = protocol.InvalidByteCount + h.min = size + h.lastProbeWasLost = false + // remove all values smaller than size from the lost array + var j int + for i, v := range h.lost { + if size < v { + j = i + break + } + } + if j > 0 { + for i := range len(h.lost) { + if i+j < len(h.lost) { + h.lost[i] = h.lost[i+j] + } else { + h.lost[i] = protocol.InvalidByteCount + } + } + } + if h.qlogger != nil { + h.qlogger.RecordEvent(qlog.MTUUpdated{ + Value: int(size), + Done: h.done(), + }) + } +} + +func (h *mtuFinderAckHandler) OnLost(wire.Frame) { + if h.generation != h.mtuFinder.generation { + // probe sent before reset received + return + } + size := h.inFlight + if size == protocol.InvalidByteCount { + panic("OnLost callback called although there's no MTU probe packet in flight") + } + h.lastProbeWasLost = true + h.inFlight = protocol.InvalidByteCount + for i, v := range h.lost { + if size < v { + copy(h.lost[i+1:], h.lost[i:]) + h.lost[i] = size + break + } + } +} diff --git a/third_party/quic-go/mtu_discoverer_test.go b/third_party/quic-go/mtu_discoverer_test.go new file mode 100644 index 0000000..c2fc990 --- /dev/null +++ b/third_party/quic-go/mtu_discoverer_test.go @@ -0,0 +1,242 @@ +package quic + +import ( + "fmt" + "math/rand/v2" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/require" +) + +func TestMTUDiscovererTiming(t *testing.T) { + const rtt = 100 * time.Millisecond + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + d := newMTUDiscoverer(rttStats, 1000, 2000, nil) + + now := monotime.Now() + require.False(t, d.ShouldSendProbe(now)) + d.Start(now) + require.False(t, d.ShouldSendProbe(now)) + require.False(t, d.ShouldSendProbe(now.Add(rtt*9/2))) + now = now.Add(5 * rtt) + require.True(t, d.ShouldSendProbe(now)) + + // only a single outstanding probe packet is permitted + ping, _ := d.GetPing(now) + require.False(t, d.ShouldSendProbe(now)) + now = now.Add(5 * rtt) + require.False(t, d.ShouldSendProbe(now)) + ping.Handler.OnLost(ping.Frame) + require.True(t, d.ShouldSendProbe(now)) +} + +func TestMTUDiscovererAckAndLoss(t *testing.T) { + const rtt = 200 * time.Millisecond + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + d := newMTUDiscoverer(rttStats, 1000, 2000, nil) + now := monotime.Now() + ping, size := d.GetPing(now) + require.Equal(t, protocol.ByteCount(1500), size) + // the MTU is reduced if the frame is lost + ping.Handler.OnLost(ping.Frame) + require.Equal(t, protocol.ByteCount(1000), d.CurrentSize()) // no change to the MTU yet + + now = now.Add(5 * rtt) + require.True(t, d.ShouldSendProbe(now)) + ping, size = d.GetPing(now) + require.Equal(t, protocol.ByteCount(1250), size) + ping.Handler.OnAcked(ping.Frame) + require.Equal(t, protocol.ByteCount(1250), d.CurrentSize()) // the MTU is increased + + // Even though the 1500 byte MTU probe packet was lost, we try again with a higher MTU. + // This protects against regular (non-MTU-related) packet loss. + now = now.Add(5 * rtt) + require.True(t, d.ShouldSendProbe(now)) + ping, size = d.GetPing(now) + require.Greater(t, size, protocol.ByteCount(1500)) + ping.Handler.OnAcked(ping.Frame) + require.Equal(t, size, d.CurrentSize()) + + // We continue probing until the MTU is close to the maximum. + var steps int + oldSize := size + now = now.Add(5 * rtt) + for d.ShouldSendProbe(now) { + ping, size = d.GetPing(now) + require.Greater(t, size, oldSize) + oldSize = size + ping.Handler.OnAcked(ping.Frame) + steps++ + require.Less(t, steps, 10) + now = now.Add(5 * rtt) + } + require.Less(t, 2000-maxMTUDiff, size) +} + +func TestMTUDiscovererMTUDiscovery(t *testing.T) { + for i := range 5 { + t.Run(fmt.Sprintf("test %d", i), func(t *testing.T) { + testMTUDiscovererMTUDiscovery(t) + }) + } +} + +func testMTUDiscovererMTUDiscovery(t *testing.T) { + const rtt = 100 * time.Millisecond + const startMTU protocol.ByteCount = 1000 + + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(rtt, 0) + + maxMTU := protocol.ByteCount(rand.IntN(int(3000-startMTU))) + startMTU + 1 + var eventRecorder events.Recorder + d := newMTUDiscoverer(rttStats, startMTU, maxMTU, &eventRecorder) + now := monotime.Now() + d.Start(now) + realMTU := protocol.ByteCount(rand.IntN(int(maxMTU-startMTU))) + startMTU + t.Logf("MTU: %d, max: %d", realMTU, maxMTU) + now = now.Add(mtuProbeDelay * rtt) + var probes []protocol.ByteCount + for d.ShouldSendProbe(now) { + require.Less(t, len(probes), 25, fmt.Sprintf("too many iterations: %v", probes)) + ping, size := d.GetPing(now) + probes = append(probes, size) + if size <= realMTU { + ping.Handler.OnAcked(ping.Frame) + } else { + ping.Handler.OnLost(ping.Frame) + } + now = now.Add(mtuProbeDelay * rtt) + } + currentMTU := d.CurrentSize() + diff := realMTU - currentMTU + require.GreaterOrEqual(t, diff, protocol.ByteCount(0)) + if maxMTU > currentMTU+maxMTU { + events := eventRecorder.Events(qlog.MTUUpdated{}) + require.NotEmpty(t, events) + require.Equal(t, qlog.MTUUpdated{Value: int(currentMTU), Done: true}, events[0]) + } + t.Logf("MTU discovered: %d (diff: %d)", currentMTU, diff) + t.Logf("probes sent (%d): %v", len(probes), probes) + require.LessOrEqual(t, diff, maxMTUDiff) +} + +func TestMTUDiscovererWithRandomLoss(t *testing.T) { + for i := range 5 { + t.Run(fmt.Sprintf("test %d", i), func(t *testing.T) { + testMTUDiscovererWithRandomLoss(t) + }) + } +} + +func testMTUDiscovererWithRandomLoss(t *testing.T) { + const rtt = 100 * time.Millisecond + const startMTU protocol.ByteCount = 1000 + const maxRandomLoss = maxLostMTUProbes - 1 + + rttStats := utils.NewRTTStats() + rttStats.SetInitialRTT(rtt) + require.Equal(t, rtt, rttStats.SmoothedRTT()) + + maxMTU := protocol.ByteCount(rand.IntN(int(3000-startMTU))) + startMTU + 1 + var eventRecorder events.Recorder + d := newMTUDiscoverer(rttStats, startMTU, maxMTU, &eventRecorder) + d.Start(monotime.Now()) + now := monotime.Now() + realMTU := protocol.ByteCount(rand.IntN(int(maxMTU-startMTU))) + startMTU + t.Logf("MTU: %d, max: %d", realMTU, maxMTU) + now = now.Add(mtuProbeDelay * rtt) + var probes, randomLosses []protocol.ByteCount + + for d.ShouldSendProbe(now) { + require.Less(t, len(probes), 32, fmt.Sprintf("too many iterations: %v", probes)) + ping, size := d.GetPing(now) + probes = append(probes, size) + packetFits := size <= realMTU + var acked bool + if packetFits { + randomLoss := rand.IntN(maxLostMTUProbes) == 0 && len(randomLosses) < maxRandomLoss + if randomLoss { + randomLosses = append(randomLosses, size) + } else { + ping.Handler.OnAcked(ping.Frame) + acked = true + } + } + if !acked { + ping.Handler.OnLost(ping.Frame) + } + now = now.Add(mtuProbeDelay * rtt) + } + + currentMTU := d.CurrentSize() + diff := realMTU - currentMTU + require.GreaterOrEqual(t, diff, protocol.ByteCount(0)) + if maxMTU > currentMTU+maxMTU { + events := eventRecorder.Events(qlog.MTUUpdated{}) + require.NotEmpty(t, events) + require.Equal(t, qlog.MTUUpdated{Value: int(currentMTU), Done: true}, events[0]) + } + t.Logf("MTU discovered with random losses %v: %d (diff: %d)", randomLosses, currentMTU, diff) + t.Logf("probes sent (%d): %v", len(probes), probes) + require.LessOrEqual(t, diff, maxMTUDiff) +} + +func TestMTUDiscovererReset(t *testing.T) { + t.Run("probe on old path acknowledged", func(t *testing.T) { + testMTUDiscovererReset(t, true) + }) + t.Run("probe on old path lost", func(t *testing.T) { + testMTUDiscovererReset(t, false) + }) +} + +func testMTUDiscovererReset(t *testing.T, ackLastProbe bool) { + const startMTU protocol.ByteCount = 1000 + const maxMTU = 1400 + const rtt = 100 * time.Millisecond + + rttStats := utils.NewRTTStats() + rttStats.SetInitialRTT(rtt) + + now := monotime.Now() + d := newMTUDiscoverer(rttStats, startMTU, maxMTU, nil) + d.Start(now) + + ping, _ := d.GetPing(now.Add(5 * rtt)) + ping.Handler.OnAcked(ping.Frame) + require.Greater(t, d.CurrentSize(), startMTU) + now = now.Add(5 * rtt) + + // send another probe packet, but neither acknowledge nor lose it before resetting + ping, _ = d.GetPing(now.Add(5 * rtt)) + now = now.Add(2 * rtt) // advance the timer by an arbitrary amount + + const newStartMTU protocol.ByteCount = 900 + const newMaxMTU = 1500 + d.Reset(now, newStartMTU, newMaxMTU) + require.Equal(t, d.CurrentSize(), newStartMTU) + + // Now acknowledge / lose the probe packet. + // This should be ignored, since it's on the old path. + if ackLastProbe { + ping.Handler.OnAcked(ping.Frame) + } else { + ping.Handler.OnLost(ping.Frame) + } + + // the MTU should not have changed + require.Equal(t, d.CurrentSize(), newStartMTU) + // the next probe should be sent after 5 RTTs + require.False(t, d.ShouldSendProbe(now.Add(5*rtt).Add(-time.Microsecond))) + require.True(t, d.ShouldSendProbe(now.Add(5*rtt))) +} diff --git a/third_party/quic-go/oss-fuzz.sh b/third_party/quic-go/oss-fuzz.sh new file mode 100644 index 0000000..f51c637 --- /dev/null +++ b/third_party/quic-go/oss-fuzz.sh @@ -0,0 +1,52 @@ +#!/bin/bash + +set -euo pipefail + +echo "Build date (UTC): $(date -u '+%Y-%m-%dT%H:%M:%SZ')" + +go version +go env + +# fuzz qpack +cd $GOPATH/src/github.com/quic-go/qpack +git log -1 --format='qpack revision: %H (%cI) %s' +compile_native_go_fuzzer_v2 github.com/quic-go/qpack FuzzDecode qpack_decode_fuzzer + +# fuzz quic-go +cd $GOPATH/src/github.com/apernet/quic-go/ +git log -1 --format='quic-go revision: %H (%cI) %s' + +build_native_go_fuzzer() { + local pkg=$1 + local fuzz=$2 + local name=$3 + local corpus_dir="${WORK:-/tmp}/quic-go-seed-corpus/$name" + local corpus_zip="$OUT/${name}_seed_corpus.zip" + + # FUZZ_CORPUS_DIR makes go-ossfuzz-seeds write each f.Add seed as a raw + # libFuzzer corpus file. OSS-Fuzz picks up _seed_corpus.zip from + # $OUT and unpacks it next to the fuzzer binary. + rm -rf "$corpus_dir" + mkdir -p "$corpus_dir" + FUZZ_CORPUS_DIR="$corpus_dir" go test "$pkg" -run "^${fuzz}$" -count=1 -v + + rm -f "$corpus_zip" + corpus_files=$(find "$corpus_dir" -type f | wc -l) + echo "$name: generated $corpus_files corpus files" + if [[ "$corpus_files" -gt 0 ]]; then + (cd "$corpus_dir" && zip -q -r "$corpus_zip" .) + fi + + compile_native_go_fuzzer_v2 "$pkg" "$fuzz" "$name" +} + +build_native_go_fuzzer github.com/apernet/quic-go/internal/wire FuzzFrames frame_fuzzer_v2 +build_native_go_fuzzer github.com/apernet/quic-go/internal/wire FuzzTransportParameters transportparameter_fuzzer_v2 +build_native_go_fuzzer github.com/apernet/quic-go/http3 FuzzFrameParser http3_frame_fuzzer +build_native_go_fuzzer github.com/apernet/quic-go/internal/wire FuzzHeaderParser header_fuzzer_v2 +build_native_go_fuzzer github.com/apernet/quic-go/internal/handshake FuzzHandshake handshake_fuzzer_v2 +build_native_go_fuzzer github.com/apernet/quic-go FuzzFrameSorter frame_sorter_fuzzer +build_native_go_fuzzer github.com/apernet/quic-go/http3 FuzzHeaderParsing http3_header_parsing_fuzzer + +# for debugging +ls -al $OUT diff --git a/third_party/quic-go/packet_packer.go b/third_party/quic-go/packet_packer.go new file mode 100644 index 0000000..739071b --- /dev/null +++ b/third_party/quic-go/packet_packer.go @@ -0,0 +1,1111 @@ +package quic + +import ( + crand "crypto/rand" + "encoding/binary" + "errors" + "fmt" + "math/rand/v2" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +var errNothingToPack = errors.New("nothing to pack") + +type packer interface { + PackCoalescedPacket(onlyAck bool, maxPacketSize protocol.ByteCount, now monotime.Time, v protocol.Version) (*coalescedPacket, error) + PackAckOnlyPacket(maxPacketSize protocol.ByteCount, now monotime.Time, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) + AppendPacket(_ *packetBuffer, maxPacketSize protocol.ByteCount, now monotime.Time, v protocol.Version) (shortHeaderPacket, error) + PackPTOProbePacket(_ protocol.EncryptionLevel, _ protocol.ByteCount, addPingIfEmpty bool, now monotime.Time, v protocol.Version) (*coalescedPacket, error) + PackConnectionClose(*qerr.TransportError, protocol.ByteCount, protocol.Version) (*coalescedPacket, error) + PackApplicationClose(*qerr.ApplicationError, protocol.ByteCount, protocol.Version) (*coalescedPacket, error) + PackPathProbePacket(protocol.ConnectionID, []ackhandler.Frame, protocol.Version) (shortHeaderPacket, *packetBuffer, error) + PackMTUProbePacket(ping ackhandler.Frame, size protocol.ByteCount, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) + + SetToken([]byte) +} + +type sealer interface { + handshake.LongHeaderSealer +} + +type payload struct { + streamFrames []ackhandler.StreamFrame + frames []ackhandler.Frame + ack *wire.AckFrame + length protocol.ByteCount +} + +type longHeaderPacket struct { + header *wire.ExtendedHeader + ack *wire.AckFrame + frames []ackhandler.Frame + streamFrames []ackhandler.StreamFrame // only used for 0-RTT packets + + length protocol.ByteCount +} + +type shortHeaderPacket struct { + PacketNumber protocol.PacketNumber + Frames []ackhandler.Frame + StreamFrames []ackhandler.StreamFrame + Ack *wire.AckFrame + Length protocol.ByteCount + IsPathMTUProbePacket bool + IsPathProbePacket bool + + // used for logging + DestConnID protocol.ConnectionID + PacketNumberLen protocol.PacketNumberLen + KeyPhase protocol.KeyPhaseBit +} + +func (p *shortHeaderPacket) IsAckEliciting() bool { return ackhandler.HasAckElicitingFrames(p.Frames) } + +type coalescedPacket struct { + buffer *packetBuffer + longHdrPackets []*longHeaderPacket + shortHdrPacket *shortHeaderPacket +} + +// IsOnlyShortHeaderPacket says if this packet only contains a short header packet (and no long header packets). +func (p *coalescedPacket) IsOnlyShortHeaderPacket() bool { + return len(p.longHdrPackets) == 0 && p.shortHdrPacket != nil +} + +func (p *longHeaderPacket) EncryptionLevel() protocol.EncryptionLevel { + //nolint:exhaustive // Will never be called for Retry packets (and they don't have encrypted data). + switch p.header.Type { + case protocol.PacketTypeInitial: + return protocol.EncryptionInitial + case protocol.PacketTypeHandshake: + return protocol.EncryptionHandshake + case protocol.PacketType0RTT: + return protocol.Encryption0RTT + default: + panic("can't determine encryption level") + } +} + +func (p *longHeaderPacket) IsAckEliciting() bool { return ackhandler.HasAckElicitingFrames(p.frames) } + +type packetNumberManager interface { + PeekPacketNumber(protocol.EncryptionLevel) (protocol.PacketNumber, protocol.PacketNumberLen) + PopPacketNumber(protocol.EncryptionLevel) protocol.PacketNumber + SetLastDatagramPadding(protocol.ByteCount) +} + +type sealingManager interface { + GetInitialSealer() (handshake.LongHeaderSealer, error) + GetHandshakeSealer() (handshake.LongHeaderSealer, error) + Get0RTTSealer() (handshake.LongHeaderSealer, error) + Get1RTTSealer() (handshake.ShortHeaderSealer, error) +} + +type frameSource interface { + HasData() bool + Append([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount, monotime.Time, protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) +} + +type ackFrameSource interface { + GetAckFrame(_ protocol.EncryptionLevel, now monotime.Time, onlyIfQueued bool) *wire.AckFrame +} + +type packetPacker struct { + srcConnID protocol.ConnectionID + getDestConnID func() protocol.ConnectionID + + perspective protocol.Perspective + cryptoSetup sealingManager + + initialStream *initialCryptoStream + handshakeStream *cryptoStream + + token []byte + + pnManager packetNumberManager + framer frameSource + acks ackFrameSource + datagramQueue *datagramQueue + retransmissionQueue *retransmissionQueue + rand rand.Rand + + numNonAckElicitingAcks int + + peekTimes int + + // chaosProtection applies Chrome's chaos protection to Initial packets. + // See appendChaosProtectedPayload. + chaosProtection bool +} + +const DatagramFrameMaxPeekTimes = 10 + +var _ packer = &packetPacker{} + +func newPacketPacker( + srcConnID protocol.ConnectionID, + getDestConnID func() protocol.ConnectionID, + initialStream *initialCryptoStream, + handshakeStream *cryptoStream, + packetNumberManager packetNumberManager, + retransmissionQueue *retransmissionQueue, + cryptoSetup sealingManager, + framer frameSource, + acks ackFrameSource, + datagramQueue *datagramQueue, + perspective protocol.Perspective, + chaosProtection bool, +) *packetPacker { + var b [16]byte + _, _ = crand.Read(b[:]) + + return &packetPacker{ + chaosProtection: chaosProtection, + cryptoSetup: cryptoSetup, + getDestConnID: getDestConnID, + srcConnID: srcConnID, + initialStream: initialStream, + handshakeStream: handshakeStream, + retransmissionQueue: retransmissionQueue, + datagramQueue: datagramQueue, + perspective: perspective, + framer: framer, + acks: acks, + rand: *rand.New(rand.NewPCG(binary.BigEndian.Uint64(b[:8]), binary.BigEndian.Uint64(b[8:]))), + pnManager: packetNumberManager, + } +} + +// noCoalescing reports whether the packet being packed has to go out in a +// datagram of its own. The imitated client never coalesces: every datagram it +// sends carries exactly one encryption level, so its Initial ACK travels alone, +// padded, and the Handshake flight follows in the next datagram. +func (p *packetPacker) noCoalescing() bool { + return p.chaosProtection && p.perspective == protocol.PerspectiveClient +} + +// splitAckFromCrypto reports whether an ACK should be sent without the CRYPTO +// data that would otherwise ride along in the same packet. See the call site. +func (p *packetPacker) splitAckFromCrypto(encLevel protocol.EncryptionLevel) bool { + return p.chaosProtection && + encLevel == protocol.EncryptionHandshake && + p.perspective == protocol.PerspectiveClient +} + +// PackConnectionClose packs a packet that closes the connection with a transport error. +func (p *packetPacker) PackConnectionClose(e *qerr.TransportError, maxPacketSize protocol.ByteCount, v protocol.Version) (*coalescedPacket, error) { + var reason string + // don't send details of crypto errors + if !e.ErrorCode.IsCryptoError() { + reason = e.ErrorMessage + } + return p.packConnectionClose(false, uint64(e.ErrorCode), e.FrameType, reason, maxPacketSize, v) +} + +// PackApplicationClose packs a packet that closes the connection with an application error. +func (p *packetPacker) PackApplicationClose(e *qerr.ApplicationError, maxPacketSize protocol.ByteCount, v protocol.Version) (*coalescedPacket, error) { + return p.packConnectionClose(true, uint64(e.ErrorCode), 0, e.ErrorMessage, maxPacketSize, v) +} + +func (p *packetPacker) packConnectionClose( + isApplicationError bool, + errorCode uint64, + frameType uint64, + reason string, + maxPacketSize protocol.ByteCount, + v protocol.Version, +) (*coalescedPacket, error) { + var sealers [4]sealer + var hdrs [3]*wire.ExtendedHeader + var payloads [4]payload + var size protocol.ByteCount + var connID protocol.ConnectionID + var oneRTTPacketNumber protocol.PacketNumber + var oneRTTPacketNumberLen protocol.PacketNumberLen + var keyPhase protocol.KeyPhaseBit // only set for 1-RTT + var numLongHdrPackets uint8 + encLevels := [4]protocol.EncryptionLevel{protocol.EncryptionInitial, protocol.EncryptionHandshake, protocol.Encryption0RTT, protocol.Encryption1RTT} + for i, encLevel := range encLevels { + if p.perspective == protocol.PerspectiveServer && encLevel == protocol.Encryption0RTT { + continue + } + ccf := &wire.ConnectionCloseFrame{ + IsApplicationError: isApplicationError, + ErrorCode: errorCode, + FrameType: frameType, + ReasonPhrase: reason, + } + // don't send application errors in Initial or Handshake packets + if isApplicationError && (encLevel == protocol.EncryptionInitial || encLevel == protocol.EncryptionHandshake) { + ccf.IsApplicationError = false + ccf.ErrorCode = uint64(qerr.ApplicationErrorErrorCode) + ccf.ReasonPhrase = "" + } + pl := payload{ + frames: []ackhandler.Frame{{Frame: ccf}}, + length: ccf.Length(v), + } + + var sealer sealer + var err error + switch encLevel { + case protocol.EncryptionInitial: + sealer, err = p.cryptoSetup.GetInitialSealer() + case protocol.EncryptionHandshake: + sealer, err = p.cryptoSetup.GetHandshakeSealer() + case protocol.Encryption0RTT: + sealer, err = p.cryptoSetup.Get0RTTSealer() + case protocol.Encryption1RTT: + var s handshake.ShortHeaderSealer + s, err = p.cryptoSetup.Get1RTTSealer() + if err == nil { + keyPhase = s.KeyPhase() + } + sealer = s + } + if err == handshake.ErrKeysNotYetAvailable || err == handshake.ErrKeysDropped { + continue + } + if err != nil { + return nil, err + } + sealers[i] = sealer + var hdr *wire.ExtendedHeader + if encLevel == protocol.Encryption1RTT { + connID = p.getDestConnID() + oneRTTPacketNumber, oneRTTPacketNumberLen = p.pnManager.PeekPacketNumber(protocol.Encryption1RTT) + size += p.shortHeaderPacketLength(connID, oneRTTPacketNumberLen, pl) + } else { + hdr = p.getLongHeader(encLevel, v) + hdrs[i] = hdr + size += p.longHeaderPacketLength(hdr, pl, v) + protocol.ByteCount(sealer.Overhead()) + numLongHdrPackets++ + } + payloads[i] = pl + } + buffer := getPacketBuffer() + packet := &coalescedPacket{ + buffer: buffer, + longHdrPackets: make([]*longHeaderPacket, 0, numLongHdrPackets), + } + for i, encLevel := range encLevels { + if sealers[i] == nil { + continue + } + if encLevel == protocol.Encryption1RTT { + shp, err := p.appendShortHeaderPacket(buffer, connID, oneRTTPacketNumber, oneRTTPacketNumberLen, keyPhase, payloads[i], 0, maxPacketSize, sealers[i], false, v) + if err != nil { + return nil, err + } + packet.shortHdrPacket = &shp + } else { + var paddingLen protocol.ByteCount + if encLevel == protocol.EncryptionInitial { + paddingLen = p.initialPaddingLen(payloads[i].frames, size, maxPacketSize) + } + longHdrPacket, err := p.appendLongHeaderPacket(buffer, hdrs[i], payloads[i], paddingLen, encLevel, sealers[i], v) + if err != nil { + buffer.Release() + return nil, err + } + packet.longHdrPackets = append(packet.longHdrPackets, longHdrPacket) + } + } + return packet, nil +} + +// longHeaderPacketLength calculates the length of a serialized long header packet. +// It takes into account that packets that have a tiny payload need to be padded, +// such that len(payload) + packet number len >= 4 + AEAD overhead +func (p *packetPacker) longHeaderPacketLength(hdr *wire.ExtendedHeader, pl payload, v protocol.Version) protocol.ByteCount { + var paddingLen protocol.ByteCount + pnLen := protocol.ByteCount(hdr.PacketNumberLen) + if pl.length < 4-pnLen { + paddingLen = 4 - pnLen - pl.length + } + return hdr.GetLength(v) + pl.length + paddingLen +} + +// shortHeaderPacketLength calculates the length of a serialized short header packet. +// It takes into account that packets that have a tiny payload need to be padded, +// such that len(payload) + packet number len >= 4 + AEAD overhead +func (p *packetPacker) shortHeaderPacketLength(connID protocol.ConnectionID, pnLen protocol.PacketNumberLen, pl payload) protocol.ByteCount { + var paddingLen protocol.ByteCount + if pl.length < 4-protocol.ByteCount(pnLen) { + paddingLen = 4 - protocol.ByteCount(pnLen) - pl.length + } + return wire.ShortHeaderLen(connID, pnLen) + pl.length + paddingLen +} + +// size is the expected size of the packet, if no padding was applied. +func (p *packetPacker) initialPaddingLen(frames []ackhandler.Frame, currentSize, maxPacketSize protocol.ByteCount) protocol.ByteCount { + // For the server, only ack-eliciting Initial packets need to be padded. + if p.perspective == protocol.PerspectiveServer && !ackhandler.HasAckElicitingFrames(frames) { + return 0 + } + if currentSize >= maxPacketSize { + return 0 + } + return maxPacketSize - currentSize +} + +// PackCoalescedPacket packs a new packet. +// It packs an Initial / Handshake if there is data to send in these packet number spaces. +// It should only be called before the handshake is confirmed. +func (p *packetPacker) PackCoalescedPacket(onlyAck bool, maxSize protocol.ByteCount, now monotime.Time, v protocol.Version) (*coalescedPacket, error) { + var ( + initialHdr, handshakeHdr, zeroRTTHdr *wire.ExtendedHeader + initialPayload, handshakePayload, zeroRTTPayload, oneRTTPayload payload + oneRTTPacketNumber protocol.PacketNumber + oneRTTPacketNumberLen protocol.PacketNumberLen + ) + // Try packing an Initial packet. + initialSealer, err := p.cryptoSetup.GetInitialSealer() + if err != nil && err != handshake.ErrKeysDropped { + return nil, err + } + var size protocol.ByteCount + if initialSealer != nil { + initialHdr, initialPayload = p.maybeGetCryptoPacket( + maxSize-protocol.ByteCount(initialSealer.Overhead()), + protocol.EncryptionInitial, + now, + false, + onlyAck, + v, + ) + if initialPayload.length > 0 { + size += p.longHeaderPacketLength(initialHdr, initialPayload, v) + protocol.ByteCount(initialSealer.Overhead()) + } + } + + // Add a Handshake packet. + var handshakeSealer sealer + appendHandshake := (onlyAck && size == 0) || (!onlyAck && size < maxSize-protocol.MinCoalescedPacketSize) + if p.noCoalescing() && size > 0 { + appendHandshake = false + } + if appendHandshake { + var err error + handshakeSealer, err = p.cryptoSetup.GetHandshakeSealer() + if err != nil && err != handshake.ErrKeysDropped && err != handshake.ErrKeysNotYetAvailable { + return nil, err + } + if handshakeSealer != nil { + handshakeHdr, handshakePayload = p.maybeGetCryptoPacket( + maxSize-size-protocol.ByteCount(handshakeSealer.Overhead()), + protocol.EncryptionHandshake, + now, + false, + onlyAck, + v, + ) + if handshakePayload.length > 0 { + s := p.longHeaderPacketLength(handshakeHdr, handshakePayload, v) + protocol.ByteCount(handshakeSealer.Overhead()) + size += s + } + } + } + + // Add a 0-RTT / 1-RTT packet. + // + // Application data is never coalesced with a long header packet during the + // handshake: those go out in datagrams of their own, with the first 1-RTT + // packet following separately. quic-go coalesces as soon as 1-RTT keys exist, + // which inflates the handshake datagrams. + var zeroRTTSealer sealer + var oneRTTSealer handshake.ShortHeaderSealer + var connID protocol.ConnectionID + var kp protocol.KeyPhaseBit + appendAppData := (onlyAck && size == 0) || (!onlyAck && size < maxSize-protocol.MinCoalescedPacketSize) + if p.noCoalescing() && size > 0 { + appendAppData = false + } + if appendAppData { + var err error + oneRTTSealer, err = p.cryptoSetup.Get1RTTSealer() + if err != nil && err != handshake.ErrKeysDropped && err != handshake.ErrKeysNotYetAvailable { + return nil, err + } + if err == nil { // 1-RTT + kp = oneRTTSealer.KeyPhase() + connID = p.getDestConnID() + oneRTTPacketNumber, oneRTTPacketNumberLen = p.pnManager.PeekPacketNumber(protocol.Encryption1RTT) + hdrLen := wire.ShortHeaderLen(connID, oneRTTPacketNumberLen) + oneRTTPayload = p.maybeGetShortHeaderPacket(oneRTTSealer, hdrLen, maxSize-size, onlyAck, now, v) + if oneRTTPayload.length > 0 { + size += p.shortHeaderPacketLength(connID, oneRTTPacketNumberLen, oneRTTPayload) + protocol.ByteCount(oneRTTSealer.Overhead()) + } + } else if p.perspective == protocol.PerspectiveClient && !onlyAck { // 0-RTT packets can't contain ACK frames + var err error + zeroRTTSealer, err = p.cryptoSetup.Get0RTTSealer() + if err != nil && err != handshake.ErrKeysDropped && err != handshake.ErrKeysNotYetAvailable { + return nil, err + } + if zeroRTTSealer != nil { + zeroRTTHdr, zeroRTTPayload = p.maybeGetAppDataPacketFor0RTT(zeroRTTSealer, maxSize-size, now, v) + if zeroRTTPayload.length > 0 { + size += p.longHeaderPacketLength(zeroRTTHdr, zeroRTTPayload, v) + protocol.ByteCount(zeroRTTSealer.Overhead()) + } + } + } + } + + if initialPayload.length == 0 && handshakePayload.length == 0 && zeroRTTPayload.length == 0 && oneRTTPayload.length == 0 { + return nil, nil + } + + buffer := getPacketBuffer() + packet := &coalescedPacket{ + buffer: buffer, + longHdrPackets: make([]*longHeaderPacket, 0, 3), + } + var padding protocol.ByteCount + if initialPayload.length > 0 { + padding = p.initialPaddingLen(initialPayload.frames, size, maxSize) + cont, err := p.appendLongHeaderPacket(buffer, initialHdr, initialPayload, padding, protocol.EncryptionInitial, initialSealer, v) + if err != nil { + buffer.Release() + return nil, err + } + packet.longHdrPackets = append(packet.longHdrPackets, cont) + } + p.pnManager.SetLastDatagramPadding(padding) + if handshakePayload.length > 0 { + cont, err := p.appendLongHeaderPacket(buffer, handshakeHdr, handshakePayload, 0, protocol.EncryptionHandshake, handshakeSealer, v) + if err != nil { + buffer.Release() + return nil, err + } + packet.longHdrPackets = append(packet.longHdrPackets, cont) + } + if zeroRTTPayload.length > 0 { + longHdrPacket, err := p.appendLongHeaderPacket(buffer, zeroRTTHdr, zeroRTTPayload, 0, protocol.Encryption0RTT, zeroRTTSealer, v) + if err != nil { + buffer.Release() + return nil, err + } + packet.longHdrPackets = append(packet.longHdrPackets, longHdrPacket) + } else if oneRTTPayload.length > 0 { + shp, err := p.appendShortHeaderPacket(buffer, connID, oneRTTPacketNumber, oneRTTPacketNumberLen, kp, oneRTTPayload, 0, maxSize, oneRTTSealer, false, v) + if err != nil { + buffer.Release() + return nil, err + } + packet.shortHdrPacket = &shp + } + return packet, nil +} + +// PackAckOnlyPacket packs a packet containing only an ACK in the application data packet number space. +// It should be called after the handshake is confirmed. +func (p *packetPacker) PackAckOnlyPacket(maxSize protocol.ByteCount, now monotime.Time, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + buf := getPacketBuffer() + packet, err := p.appendPacket(buf, true, maxSize, now, v) + return packet, buf, err +} + +// AppendPacket packs a packet in the application data packet number space. +// It should be called after the handshake is confirmed. +func (p *packetPacker) AppendPacket(buf *packetBuffer, maxSize protocol.ByteCount, now monotime.Time, v protocol.Version) (shortHeaderPacket, error) { + return p.appendPacket(buf, false, maxSize, now, v) +} + +func (p *packetPacker) appendPacket( + buf *packetBuffer, + onlyAck bool, + maxPacketSize protocol.ByteCount, + now monotime.Time, + v protocol.Version, +) (shortHeaderPacket, error) { + sealer, err := p.cryptoSetup.Get1RTTSealer() + if err != nil { + return shortHeaderPacket{}, err + } + pn, pnLen := p.pnManager.PeekPacketNumber(protocol.Encryption1RTT) + connID := p.getDestConnID() + hdrLen := wire.ShortHeaderLen(connID, pnLen) + pl := p.maybeGetShortHeaderPacket(sealer, hdrLen, maxPacketSize, onlyAck, now, v) + if pl.length == 0 { + return shortHeaderPacket{}, errNothingToPack + } + kp := sealer.KeyPhase() + + return p.appendShortHeaderPacket(buf, connID, pn, pnLen, kp, pl, 0, maxPacketSize, sealer, false, v) +} + +func (p *packetPacker) maybeGetCryptoPacket( + maxPacketSize protocol.ByteCount, + encLevel protocol.EncryptionLevel, + now monotime.Time, + addPingIfEmpty bool, + onlyAck bool, + v protocol.Version, +) (*wire.ExtendedHeader, payload) { + if onlyAck { + if ack := p.acks.GetAckFrame(encLevel, now, true); ack != nil { + hdr := p.getLongHeader(encLevel, v) + maxPacketSize -= hdr.GetLength(v) + ack.Truncate(maxPacketSize, v) + return hdr, payload{ack: ack, length: ack.Length(v)} + } + return nil, payload{length: 0} + } + + var hasCryptoData func() bool + var popCryptoFrame func(maxLen protocol.ByteCount) *wire.CryptoFrame + var pendingCryptoLen func() protocol.ByteCount + var cryptoWriteOffset func() protocol.ByteCount + var popCryptoFrameTail func(dataLen protocol.ByteCount) *wire.CryptoFrame + //nolint:exhaustive // Initial and Handshake are the only two encryption levels here. + switch encLevel { + case protocol.EncryptionInitial: + hasCryptoData = p.initialStream.HasData + popCryptoFrame = p.initialStream.PopCryptoFrame + pendingCryptoLen = p.initialStream.PendingLen + cryptoWriteOffset = p.initialStream.WriteOffset + popCryptoFrameTail = p.initialStream.PopCryptoFrameTail + case protocol.EncryptionHandshake: + hasCryptoData = p.handshakeStream.HasData + popCryptoFrame = p.handshakeStream.PopCryptoFrame + pendingCryptoLen = p.handshakeStream.PendingLen + cryptoWriteOffset = p.handshakeStream.WriteOffset + } + handler := p.retransmissionQueue.AckHandler(encLevel) + hasRetransmission := p.retransmissionQueue.HasData(encLevel) + + ack := p.acks.GetAckFrame(encLevel, now, !hasRetransmission && !hasCryptoData()) + var pl payload + if !hasCryptoData() && !hasRetransmission && ack == nil { + if !addPingIfEmpty { + // nothing to send + return nil, payload{} + } + ping := &wire.PingFrame{} + pl.frames = append(pl.frames, ackhandler.Frame{Frame: ping, Handler: emptyHandler{}}) + pl.length += ping.Length(v) + } + + hdr := p.getLongHeader(encLevel, v) + maxPacketSize -= hdr.GetLength(v) + + if ack != nil { + ack.Truncate(maxPacketSize, v) + pl.ack = ack + pl.length = ack.Length(v) + maxPacketSize -= pl.length + } + // The server's Handshake flight is acknowledged in a datagram of its own, with + // the Finished following in the next one; quic-go packs both into a single + // packet. Returning the ACK alone leaves the CRYPTO for the next datagram, + // which the send loop packs immediately afterwards. + if !hasRetransmission && p.splitAckFromCrypto(encLevel) && pl.ack != nil && hasCryptoData() { + return hdr, pl + } + if hasRetransmission { + for { + frame := p.retransmissionQueue.GetFrame(encLevel, maxPacketSize, v) + if frame == nil { + break + } + pl.frames = append(pl.frames, ackhandler.Frame{ + Frame: frame, + Handler: p.retransmissionQueue.AckHandler(encLevel), + }) + frameLen := frame.Length(v) + pl.length += frameLen + maxPacketSize -= frameLen + } + return hdr, pl + } else { + // A ClientHello too large for one Initial is not simply cut in two: the + // first packet carries its head and its tail, and the middle follows in + // later packets. See chromeCryptoSplit. + if p.chaosProtection && encLevel == protocol.EncryptionInitial && + pendingCryptoLen != nil && popCryptoFrameTail != nil && cryptoWriteOffset() == 0 { + first, last := chromeCryptoSplit(pendingCryptoLen(), cryptoWriteOffset(), maxPacketSize, p.rand.IntN) + if first > 0 && last > 0 { + // Take the tail before the head, or the head pop consumes it. + tail := popCryptoFrameTail(last) + head := popCryptoFrame(first + cryptoFrameHeaderLen(cryptoWriteOffset(), first)) + for _, cf := range []*wire.CryptoFrame{head, tail} { + if cf == nil { + continue + } + pl.frames = append(pl.frames, ackhandler.Frame{Frame: cf, Handler: handler}) + pl.length += cf.Length(v) + maxPacketSize -= cf.Length(v) + } + return hdr, pl + } + } + for hasCryptoData() { + cf := popCryptoFrame(maxPacketSize) + if cf == nil { + break + } + pl.frames = append(pl.frames, ackhandler.Frame{Frame: cf, Handler: handler}) + pl.length += cf.Length(v) + maxPacketSize -= cf.Length(v) + } + } + return hdr, pl +} + +func (p *packetPacker) maybeGetAppDataPacketFor0RTT(sealer sealer, maxSize protocol.ByteCount, now monotime.Time, v protocol.Version) (*wire.ExtendedHeader, payload) { + if p.perspective != protocol.PerspectiveClient { + return nil, payload{} + } + + hdr := p.getLongHeader(protocol.Encryption0RTT, v) + maxPayloadSize := maxSize - hdr.GetLength(v) - protocol.ByteCount(sealer.Overhead()) + return hdr, p.maybeGetAppDataPacket(maxPayloadSize, false, false, now, v) +} + +func (p *packetPacker) maybeGetShortHeaderPacket( + sealer handshake.ShortHeaderSealer, + hdrLen, maxPacketSize protocol.ByteCount, + onlyAck bool, + now monotime.Time, + v protocol.Version, +) payload { + maxPayloadSize := maxPacketSize - hdrLen - protocol.ByteCount(sealer.Overhead()) + return p.maybeGetAppDataPacket(maxPayloadSize, onlyAck, true, now, v) +} + +func (p *packetPacker) maybeGetAppDataPacket( + maxPayloadSize protocol.ByteCount, + onlyAck, ackAllowed bool, + now monotime.Time, + v protocol.Version, +) payload { + pl := p.composeNextPacket(maxPayloadSize, onlyAck, ackAllowed, now, v) + + // check if we have anything to send + if len(pl.frames) == 0 && len(pl.streamFrames) == 0 { + if pl.ack == nil { + return payload{} + } + // the packet only contains an ACK + if p.numNonAckElicitingAcks >= protocol.MaxNonAckElicitingAcks { + ping := &wire.PingFrame{} + pl.frames = append(pl.frames, ackhandler.Frame{Frame: ping}) + pl.length += ping.Length(v) + p.numNonAckElicitingAcks = 0 + } else { + p.numNonAckElicitingAcks++ + } + } else { + p.numNonAckElicitingAcks = 0 + } + return pl +} + +func (p *packetPacker) composeNextPacket( + maxPayloadSize protocol.ByteCount, + onlyAck, ackAllowed bool, + now monotime.Time, + v protocol.Version, +) payload { + if onlyAck { + if ack := p.acks.GetAckFrame(protocol.Encryption1RTT, now, true); ack != nil { + ack.Truncate(maxPayloadSize, v) + return payload{ack: ack, length: ack.Length(v)} + } + return payload{} + } + + hasData := p.framer.HasData() + hasRetransmission := p.retransmissionQueue.HasData(protocol.Encryption1RTT) + + var pl payload + if ackAllowed { + if ack := p.acks.GetAckFrame(protocol.Encryption1RTT, now, !hasRetransmission && !hasData); ack != nil { + ack.Truncate(maxPayloadSize, v) + pl.ack = ack + pl.length += ack.Length(v) + } + } + + if p.datagramQueue != nil { + if f := p.datagramQueue.Peek(); f != nil { + size := f.Length(v) + if size <= maxPayloadSize-pl.length { // DATAGRAM frame fits + pl.frames = append(pl.frames, ackhandler.Frame{Frame: f}) + pl.length += size + p.datagramQueue.Pop() + p.peekTimes = 0 + } else if pl.ack == nil { + // The DATAGRAM frame doesn't fit, and the packet doesn't contain an ACK. + // Discard this frame. There's no point in retrying this in the next packet, + // as it's unlikely that the available packet size will increase. + p.datagramQueue.Pop() + p.peekTimes = 0 + } + // If the DATAGRAM frame was too large and the packet contained an ACK, we'll try to send it out later. + p.peekTimes++ + if p.peekTimes > DatagramFrameMaxPeekTimes { + if p.datagramQueue.logger != nil && p.datagramQueue.logger.Debug() { + p.datagramQueue.logger.Debugf("Discarded DATAGRAM frame (%d bytes payload)", size) + } + p.datagramQueue.Pop() + p.peekTimes = 0 + } + } + } + + if pl.ack != nil && !hasData && !hasRetransmission { + return pl + } + + if hasRetransmission { + for { + remainingLen := maxPayloadSize - pl.length + if remainingLen < protocol.MinStreamFrameSize { + break + } + f := p.retransmissionQueue.GetFrame(protocol.Encryption1RTT, remainingLen, v) + if f == nil { + break + } + pl.frames = append(pl.frames, ackhandler.Frame{Frame: f, Handler: p.retransmissionQueue.AckHandler(protocol.Encryption1RTT)}) + pl.length += f.Length(v) + } + } + + if hasData { + var lengthAdded protocol.ByteCount + startLen := len(pl.frames) + pl.frames, pl.streamFrames, lengthAdded = p.framer.Append(pl.frames, pl.streamFrames, maxPayloadSize-pl.length, now, v) + pl.length += lengthAdded + // add handlers for the control frames that were added + for i := startLen; i < len(pl.frames); i++ { + if pl.frames[i].Handler != nil { + continue + } + switch pl.frames[i].Frame.(type) { + case *wire.PathChallengeFrame, *wire.PathResponseFrame: + // Path probing is currently not supported, therefore we don't need to set the OnAcked callback yet. + // PATH_CHALLENGE and PATH_RESPONSE are never retransmitted. + default: + // we might be packing a 0-RTT packet, but we need to use the 1-RTT ack handler anyway + pl.frames[i].Handler = p.retransmissionQueue.AckHandler(protocol.Encryption1RTT) + } + } + } + return pl +} + +func (p *packetPacker) PackPTOProbePacket( + encLevel protocol.EncryptionLevel, + maxPacketSize protocol.ByteCount, + addPingIfEmpty bool, + now monotime.Time, + v protocol.Version, +) (*coalescedPacket, error) { + if encLevel == protocol.Encryption1RTT { + return p.packPTOProbePacket1RTT(maxPacketSize, addPingIfEmpty, now, v) + } + + var sealer handshake.LongHeaderSealer + //nolint:exhaustive // Probe packets are never sent for 0-RTT. + switch encLevel { + case protocol.EncryptionInitial: + var err error + sealer, err = p.cryptoSetup.GetInitialSealer() + if err != nil { + return nil, err + } + case protocol.EncryptionHandshake: + var err error + sealer, err = p.cryptoSetup.GetHandshakeSealer() + if err != nil { + return nil, err + } + default: + panic("unknown encryption level") + } + hdr, pl := p.maybeGetCryptoPacket( + maxPacketSize-protocol.ByteCount(sealer.Overhead()), + encLevel, + now, + addPingIfEmpty, + false, + v, + ) + if pl.length == 0 { + return nil, nil + } + buffer := getPacketBuffer() + packet := &coalescedPacket{buffer: buffer} + size := p.longHeaderPacketLength(hdr, pl, v) + protocol.ByteCount(sealer.Overhead()) + var padding protocol.ByteCount + if encLevel == protocol.EncryptionInitial { + padding = p.initialPaddingLen(pl.frames, size, maxPacketSize) + } + longHdrPacket, err := p.appendLongHeaderPacket(buffer, hdr, pl, padding, encLevel, sealer, v) + if err != nil { + buffer.Release() + return nil, err + } + p.pnManager.SetLastDatagramPadding(padding) + packet.longHdrPackets = []*longHeaderPacket{longHdrPacket} + return packet, nil +} + +func (p *packetPacker) packPTOProbePacket1RTT(maxPacketSize protocol.ByteCount, addPingIfEmpty bool, now monotime.Time, v protocol.Version) (*coalescedPacket, error) { + s, err := p.cryptoSetup.Get1RTTSealer() + if err != nil { + return nil, err + } + kp := s.KeyPhase() + connID := p.getDestConnID() + pn, pnLen := p.pnManager.PeekPacketNumber(protocol.Encryption1RTT) + hdrLen := wire.ShortHeaderLen(connID, pnLen) + pl := p.maybeGetAppDataPacket(maxPacketSize-protocol.ByteCount(s.Overhead())-hdrLen, false, true, now, v) + if pl.length == 0 { + if !addPingIfEmpty { + return nil, nil + } + ping := &wire.PingFrame{} + pl.frames = append(pl.frames, ackhandler.Frame{Frame: ping, Handler: emptyHandler{}}) + pl.length += ping.Length(v) + } + buffer := getPacketBuffer() + packet := &coalescedPacket{buffer: buffer} + shp, err := p.appendShortHeaderPacket(buffer, connID, pn, pnLen, kp, pl, 0, maxPacketSize, s, false, v) + if err != nil { + buffer.Release() + return nil, err + } + packet.shortHdrPacket = &shp + return packet, nil +} + +func (p *packetPacker) PackMTUProbePacket(ping ackhandler.Frame, size protocol.ByteCount, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + pl := payload{ + frames: []ackhandler.Frame{ping}, + length: ping.Frame.Length(v), + } + buffer := getPacketBuffer() + s, err := p.cryptoSetup.Get1RTTSealer() + if err != nil { + return shortHeaderPacket{}, nil, err + } + connID := p.getDestConnID() + pn, pnLen := p.pnManager.PeekPacketNumber(protocol.Encryption1RTT) + padding := size - p.shortHeaderPacketLength(connID, pnLen, pl) - protocol.ByteCount(s.Overhead()) + kp := s.KeyPhase() + packet, err := p.appendShortHeaderPacket(buffer, connID, pn, pnLen, kp, pl, padding, size, s, true, v) + if err != nil { + buffer.Release() + } + return packet, buffer, err +} + +func (p *packetPacker) PackPathProbePacket(connID protocol.ConnectionID, frames []ackhandler.Frame, v protocol.Version) (shortHeaderPacket, *packetBuffer, error) { + pn, pnLen := p.pnManager.PeekPacketNumber(protocol.Encryption1RTT) + buf := getPacketBuffer() + s, err := p.cryptoSetup.Get1RTTSealer() + if err != nil { + return shortHeaderPacket{}, nil, err + } + var l protocol.ByteCount + for _, f := range frames { + l += f.Frame.Length(v) + } + payload := payload{ + frames: frames, + length: l, + } + padding := protocol.MinInitialPacketSize - p.shortHeaderPacketLength(connID, pnLen, payload) - protocol.ByteCount(s.Overhead()) + packet, err := p.appendShortHeaderPacket(buf, connID, pn, pnLen, s.KeyPhase(), payload, padding, protocol.MinInitialPacketSize, s, false, v) + if err != nil { + return shortHeaderPacket{}, nil, err + } + packet.IsPathProbePacket = true + return packet, buf, err +} + +func (p *packetPacker) getLongHeader(encLevel protocol.EncryptionLevel, v protocol.Version) *wire.ExtendedHeader { + pn, pnLen := p.pnManager.PeekPacketNumber(encLevel) + hdr := &wire.ExtendedHeader{ + PacketNumber: pn, + PacketNumberLen: pnLen, + } + hdr.Version = v + hdr.SrcConnectionID = p.srcConnID + hdr.DestConnectionID = p.getDestConnID() + + //nolint:exhaustive // 1-RTT packets are not long header packets. + switch encLevel { + case protocol.EncryptionInitial: + hdr.Type = protocol.PacketTypeInitial + hdr.Token = p.token + case protocol.EncryptionHandshake: + hdr.Type = protocol.PacketTypeHandshake + case protocol.Encryption0RTT: + hdr.Type = protocol.PacketType0RTT + } + return hdr +} + +func (p *packetPacker) appendLongHeaderPacket(buffer *packetBuffer, header *wire.ExtendedHeader, pl payload, padding protocol.ByteCount, encLevel protocol.EncryptionLevel, sealer sealer, v protocol.Version) (*longHeaderPacket, error) { + var paddingLen protocol.ByteCount + pnLen := protocol.ByteCount(header.PacketNumberLen) + if pl.length < 4-pnLen { + paddingLen = 4 - pnLen - pl.length + } + paddingLen += padding + header.Length = pnLen + protocol.ByteCount(sealer.Overhead()) + pl.length + paddingLen + + startLen := len(buffer.Data) + raw := buffer.Data[startLen:] + raw, err := header.Append(raw, v) + if err != nil { + return nil, err + } + payloadOffset := protocol.ByteCount(len(raw)) + + if p.chaosProtection && encLevel == protocol.EncryptionInitial && worthChaosProtecting(pl, paddingLen) { + raw, err = p.appendChaosProtectedPayload(raw, pl, paddingLen, v) + } else { + raw, err = p.appendPacketPayload(raw, pl, paddingLen, v) + } + if err != nil { + return nil, err + } + raw = p.encryptPacket(raw, sealer, header.PacketNumber, payloadOffset, pnLen) + buffer.Data = buffer.Data[:len(buffer.Data)+len(raw)] + + if pn := p.pnManager.PopPacketNumber(encLevel); pn != header.PacketNumber { + return nil, fmt.Errorf("packetPacker BUG: Peeked and Popped packet numbers do not match: expected %d, got %d", pn, header.PacketNumber) + } + return &longHeaderPacket{ + header: header, + ack: pl.ack, + frames: pl.frames, + streamFrames: pl.streamFrames, + length: protocol.ByteCount(len(raw)), + }, nil +} + +func (p *packetPacker) appendShortHeaderPacket( + buffer *packetBuffer, + connID protocol.ConnectionID, + pn protocol.PacketNumber, + pnLen protocol.PacketNumberLen, + kp protocol.KeyPhaseBit, + pl payload, + padding, maxPacketSize protocol.ByteCount, + sealer sealer, + isMTUProbePacket bool, + v protocol.Version, +) (shortHeaderPacket, error) { + var paddingLen protocol.ByteCount + if pl.length < 4-protocol.ByteCount(pnLen) { + paddingLen = 4 - protocol.ByteCount(pnLen) - pl.length + } + paddingLen += padding + + startLen := len(buffer.Data) + raw := buffer.Data[startLen:] + raw, err := wire.AppendShortHeader(raw, connID, pn, pnLen, kp) + if err != nil { + return shortHeaderPacket{}, err + } + payloadOffset := protocol.ByteCount(len(raw)) + + raw, err = p.appendPacketPayload(raw, pl, paddingLen, v) + if err != nil { + return shortHeaderPacket{}, err + } + if !isMTUProbePacket { + if size := protocol.ByteCount(len(raw) + sealer.Overhead()); size > maxPacketSize { + return shortHeaderPacket{}, fmt.Errorf("PacketPacker BUG: packet too large (%d bytes, allowed %d bytes)", size, maxPacketSize) + } + } + raw = p.encryptPacket(raw, sealer, pn, payloadOffset, protocol.ByteCount(pnLen)) + buffer.Data = buffer.Data[:len(buffer.Data)+len(raw)] + + if newPN := p.pnManager.PopPacketNumber(protocol.Encryption1RTT); newPN != pn { + return shortHeaderPacket{}, fmt.Errorf("packetPacker BUG: Peeked and Popped packet numbers do not match: expected %d, got %d", pn, newPN) + } + return shortHeaderPacket{ + PacketNumber: pn, + PacketNumberLen: pnLen, + KeyPhase: kp, + StreamFrames: pl.streamFrames, + Frames: pl.frames, + Ack: pl.ack, + Length: protocol.ByteCount(len(raw)), + DestConnID: connID, + IsPathMTUProbePacket: isMTUProbePacket, + }, nil +} + +// appendPacketPayload serializes the payload of a packet into the raw byte slice. +// It modifies the order of payload.frames. +func (p *packetPacker) appendPacketPayload(raw []byte, pl payload, paddingLen protocol.ByteCount, v protocol.Version) ([]byte, error) { + payloadOffset := len(raw) + if pl.ack != nil { + var err error + raw, err = pl.ack.Append(raw, v) + if err != nil { + return nil, err + } + } + if paddingLen > 0 { + raw = append(raw, make([]byte, paddingLen)...) + } + // Randomize the order of the control frames. + // This makes sure that the receiver doesn't rely on the order in which frames are packed. + if len(pl.frames) > 1 { + p.rand.Shuffle(len(pl.frames), func(i, j int) { pl.frames[i], pl.frames[j] = pl.frames[j], pl.frames[i] }) + } + for _, f := range pl.frames { + var err error + raw, err = f.Frame.Append(raw, v) + if err != nil { + return nil, err + } + } + for _, f := range pl.streamFrames { + var err error + raw, err = f.Frame.Append(raw, v) + if err != nil { + return nil, err + } + } + + if payloadSize := protocol.ByteCount(len(raw)-payloadOffset) - paddingLen; payloadSize != pl.length { + return nil, fmt.Errorf("PacketPacker BUG: payload size inconsistent (expected %d, got %d bytes)", pl.length, payloadSize) + } + return raw, nil +} + +func (p *packetPacker) encryptPacket(raw []byte, sealer sealer, pn protocol.PacketNumber, payloadOffset, pnLen protocol.ByteCount) []byte { + _ = sealer.Seal(raw[payloadOffset:payloadOffset], raw[payloadOffset:], pn, raw[:payloadOffset]) + raw = raw[:len(raw)+sealer.Overhead()] + // apply header protection + pnOffset := payloadOffset - pnLen + sealer.EncryptHeader(raw[pnOffset+4:pnOffset+4+16], &raw[0], raw[pnOffset:payloadOffset]) + return raw +} + +func (p *packetPacker) SetToken(token []byte) { + p.token = token +} + +type emptyHandler struct{} + +var _ ackhandler.FrameHandler = emptyHandler{} + +func (emptyHandler) OnAcked(wire.Frame) {} +func (emptyHandler) OnLost(wire.Frame) {} diff --git a/third_party/quic-go/packet_packer_chaos.go b/third_party/quic-go/packet_packer_chaos.go new file mode 100644 index 0000000..208a675 --- /dev/null +++ b/third_party/quic-go/packet_packer_chaos.go @@ -0,0 +1,274 @@ +package quic + +import ( + "fmt" + + "github.com/apernet/quic-go/quicvarint" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +// Bounds for the chaos protection applied to Initial packets. +const ( + // How many extra CRYPTO frames to attempt. An attempt that lands on a frame + // that cannot be split is spent doing nothing, so the resulting frame count + // is a distribution rather than a fixed target. + chaosMinAddedCryptoFrames = 2 + chaosMaxAddedCryptoFrames = 10 + + chaosMinPingFrames = 2 + chaosMaxPingFrames = 10 + + // Worst-case encoded size of a CRYPTO frame header: type byte plus a varint + // offset and a varint length. Used to bound per-packet capacity. + maxCryptoFrameHeaderLen = 8 +) + +// Length of the leading CRYPTO frame when a ClientHello is spread over several +// Initials. It covers the fixed part of the ClientHello plus the length of its +// first extension, so a receiver must reassemble to parse past it, with enough +// randomness that the cut point is not a constant. +const ( + chaosMinFirstFrameLen = 55 + chaosFirstFrameLenRandom = 32 +) + +// chromeCryptoSplit decides how a ClientHello too large for one Initial is laid +// out. It returns the length of the leading frame and of the trailing frame, +// which the first packet carries together: the head of the ClientHello and its +// very end, with the middle left for the packets that follow. +// +// Carrying the two ends first means a receiver cannot parse the ClientHello +// from the first packet alone, nor assume CRYPTO offsets arrive in order. +// Zero means the data fits in one packet and no split is needed. +// +// pending is the number of bytes queued, offset the crypto stream offset they +// start at, and available the payload space of a packet. +func chromeCryptoSplit(pending, offset, available protocol.ByteCount, rand func(int) int) (first, last protocol.ByteCount) { + if pending <= 0 || available <= 0 { + return 0, 0 + } + // Ceiling on one frame header over the whole range, used to reserve space. + minFrame := cryptoFrameHeaderLen(offset+pending-1, pending) + if available < 2*minFrame { + return 0, 0 + } + // The first packet holds two frames, later packets one, so the first has one + // header less to spend on data. + maxFirst := available - 2*minFrame + maxOther := available - minFrame + if pending <= maxFirst { + return 0, 0 // fits in a single packet + } + occupied := maxOther - maxFirst + numPackets := (pending + occupied + maxOther - 1) / maxOther // ceil + dataOther := (pending + occupied + numPackets - 1) / numPackets + if dataOther < occupied { + return 0, 0 + } + dataFirst := dataOther - occupied + + first = protocol.ByteCount(chaosMinFirstFrameLen + rand(chaosFirstFrameLenRandom)) + if pending <= first || dataFirst <= first { + return 0, 0 + } + return first, dataFirst - first +} + +// cryptoFrameHeaderLen is the encoded size of a CRYPTO frame header carrying the +// given offset and data length: type byte plus two varints. +func cryptoFrameHeaderLen(offset, dataLen protocol.ByteCount) protocol.ByteCount { + return 1 + protocol.ByteCount(quicvarint.Len(uint64(offset))) + + protocol.ByteCount(quicvarint.Len(uint64(dataLen))) +} + +// chaosItem is one element of the shuffled packet payload: either a frame to +// encode, or a run of padding bytes. +type chaosItem struct { + frame wire.Frame + paddingLen protocol.ByteCount +} + +// appendChaosProtectedPayload writes the payload the way Chromium's +// QuicChaosProtector does: the CRYPTO data is cut into several frames, PING +// frames are mixed in, the padding is broken into runs, and the lot is shuffled. +// +// This denies a DPI box the ability to parse a ClientHello out of a single +// contiguous CRYPTO frame at a predictable offset, and is a fingerprint in its +// own right, since stock quic-go never does it. +// +// The encoded size is preserved exactly: every byte spent on extra frame headers +// and PINGs comes out of the padding budget. If the budget can't cover a step we +// stop splitting rather than overrun. +func (p *packetPacker) appendChaosProtectedPayload(raw []byte, pl payload, paddingLen protocol.ByteCount, v protocol.Version) ([]byte, error) { + startLen := len(raw) + + // The ACK, if any, is written first and left out of the shuffle. Packets that + // carry an ACK and no CRYPTO data never get here, see worthChaosProtecting. + if pl.ack != nil { + var err error + if raw, err = pl.ack.Append(raw, v); err != nil { + return nil, err + } + } + + crypto, others := splitCryptoFrames(pl.frames) + budget := paddingLen + + crypto, budget = p.shredCryptoFrames(crypto, len(others), budget) + + items := make([]chaosItem, 0, len(crypto)+len(others)+2*chaosMaxPingFrames) + for _, cf := range crypto { + items = append(items, chaosItem{frame: cf}) + } + for _, f := range others { + items = append(items, chaosItem{frame: f.Frame}) + } + + // PING frames are one byte each and indistinguishable from padding to anything + // not decrypting the packet. + numPings := min( + protocol.ByteCount(chaosMinPingFrames+p.rand.IntN(chaosMaxPingFrames-chaosMinPingFrames+1)), + budget, + ) + for range numPings { + items = append(items, chaosItem{frame: &wire.PingFrame{}}) + } + budget -= numPings + + // Padding is handed out across the frames accumulated so far, before they + // are shuffled, so the run count tracks the frame count. + for _, run := range p.spreadPadding(len(items), budget) { + items = append(items, chaosItem{paddingLen: run}) + } + + p.rand.Shuffle(len(items), func(i, j int) { items[i], items[j] = items[j], items[i] }) + + for _, item := range items { + if item.frame != nil { + var err error + if raw, err = item.frame.Append(raw, v); err != nil { + return nil, err + } + continue + } + raw = append(raw, make([]byte, item.paddingLen)...) + } + + // The whole point is that the packet size is unchanged, so verify it. + if written := protocol.ByteCount(len(raw) - startLen); written != pl.length+paddingLen { + return nil, fmt.Errorf("packetPacker BUG: chaos-protected payload size inconsistent (expected %d, got %d bytes)", pl.length+paddingLen, written) + } + return raw, nil +} + +// splitCryptoFrames partitions a frame list into CRYPTO frames and everything +// else. Only CRYPTO frames can be shredded, since CRYPTO carries an explicit +// offset and may legally be split at any byte boundary. +// worthChaosProtecting reports whether a packet is one the imitated client +// scrambles. It needs CRYPTO data to shred and padding to spend on the result, +// so an Initial carrying only an ACK goes out as a plain ACK and PADDING. +func worthChaosProtecting(pl payload, paddingLen protocol.ByteCount) bool { + if paddingLen <= 0 { + return false + } + for _, f := range pl.frames { + if _, ok := f.Frame.(*wire.CryptoFrame); ok { + return true + } + } + return false +} + +func splitCryptoFrames(frames []ackhandler.Frame) ([]*wire.CryptoFrame, []ackhandler.Frame) { + crypto := make([]*wire.CryptoFrame, 0, len(frames)) + others := make([]ackhandler.Frame, 0, len(frames)) + for _, f := range frames { + if cf, ok := f.Frame.(*wire.CryptoFrame); ok { + crypto = append(crypto, cf) + continue + } + others = append(others, f) + } + return crypto, others +} + +// shredCryptoFrames makes a bounded number of attempts to split a randomly +// chosen frame, stopping early once the padding budget can no longer cover +// another frame header. Returns the new frame set and the remaining budget. +// +// The choice ranges over every frame in the packet, not just the splittable +// ones, so an attempt that lands elsewhere is simply spent. Always splitting +// the largest frame instead would produce evenly sized pieces, which is its own +// signature. +func (p *packetPacker) shredCryptoFrames(crypto []*wire.CryptoFrame, numOther int, budget protocol.ByteCount) ([]*wire.CryptoFrame, protocol.ByteCount) { + if len(crypto) == 0 { + return crypto, budget + } + // Ceiling on what one more frame header can cost, over the whole CRYPTO + // range and computed once, so the stop condition doesn't drift as we split. + lo, hi := crypto[0].Offset, protocol.ByteCount(0) + for _, cf := range crypto { + lo = min(lo, cf.Offset) + hi = max(hi, cf.Offset+protocol.ByteCount(len(cf.Data))) + } + maxOverhead := cryptoFrameHeaderLen(hi, hi-lo) + + attempts := chaosMinAddedCryptoFrames + p.rand.IntN(chaosMaxAddedCryptoFrames-chaosMinAddedCryptoFrames+1) + for range attempts { + if budget < maxOverhead { + break + } + idx := p.rand.IntN(len(crypto) + numOther) + if idx >= len(crypto) { + continue // landed on a frame we can't split + } + cf := crypto[idx] + if len(cf.Data) <= 1 { + continue + } + cut := 1 + p.rand.IntN(len(cf.Data)-1) + a := &wire.CryptoFrame{Offset: cf.Offset, Data: cf.Data[:cut]} + b := &wire.CryptoFrame{Offset: cf.Offset + protocol.ByteCount(cut), Data: cf.Data[cut:]} + + budget += cryptoFrameHeaderLen(cf.Offset, protocol.ByteCount(len(cf.Data))) + budget -= cryptoFrameHeaderLen(a.Offset, protocol.ByteCount(len(a.Data))) + budget -= cryptoFrameHeaderLen(b.Offset, protocol.ByteCount(len(b.Data))) + + crypto[idx] = a + crypto = append(crypto, b) + } + return crypto, budget +} + +// spreadPadding breaks a padding budget into runs summing to exactly the +// budget, so padding appears in several places rather than one contiguous +// block. +// +// Each of the numFrames positions in turn takes a uniformly random share of +// whatever is left, so the first runs tend to be large and later ones short or +// absent; any remainder becomes a final run. Splitting the budget evenly, or at +// sorted cut points, produces a visibly different mix of run lengths. +func (p *packetPacker) spreadPadding(numFrames int, budget protocol.ByteCount) []protocol.ByteCount { + if budget <= 0 { + return nil + } + runs := make([]protocol.ByteCount, 0, numFrames+1) + for range numFrames { + if budget <= 0 { + break + } + n := protocol.ByteCount(p.rand.IntN(int(budget) + 1)) + if n <= 0 { + continue + } + runs = append(runs, n) + budget -= n + } + if budget > 0 { + runs = append(runs, budget) + } + return runs +} diff --git a/third_party/quic-go/packet_packer_chaos_test.go b/third_party/quic-go/packet_packer_chaos_test.go new file mode 100644 index 0000000..85bc236 --- /dev/null +++ b/third_party/quic-go/packet_packer_chaos_test.go @@ -0,0 +1,292 @@ +package quic + +import ( + "testing" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/quicvarint" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +// chaosTestPayload builds a payload holding one contiguous CRYPTO frame, which +// is what the Initial packet carrying a ClientHello looks like before chaos +// protection is applied. +func chaosTestPayload(dataLen int) payload { + data := make([]byte, dataLen) + for i := range data { + data[i] = byte(i) + } + cf := &wire.CryptoFrame{Data: data} + return payload{ + frames: []ackhandler.Frame{{Frame: cf, Handler: emptyHandler{}}}, + length: cf.Length(protocol.Version1), + } +} + +// parseChaosPayload decodes an assembled payload back into CRYPTO frames, a PING +// count and a padding byte count. +// +// A chaos-protected Initial payload only ever contains PADDING (0x00), PING +// (0x01) and CRYPTO (0x06), so this walks those three directly rather than +// going through wire.FrameParser, whose ParseType silently skips padding. +func parseChaosPayload(t *testing.T, raw []byte) (crypto []*wire.CryptoFrame, pings, padding int) { + t.Helper() + for len(raw) > 0 { + switch raw[0] { + case 0x00: // PADDING + padding++ + raw = raw[1:] + case 0x01: // PING + pings++ + raw = raw[1:] + case 0x06: // CRYPTO + raw = raw[1:] + offset, n, err := quicvarint.Parse(raw) + require.NoError(t, err) + raw = raw[n:] + length, n, err := quicvarint.Parse(raw) + require.NoError(t, err) + raw = raw[n:] + require.GreaterOrEqual(t, uint64(len(raw)), length, "truncated CRYPTO frame") + crypto = append(crypto, &wire.CryptoFrame{ + Offset: protocol.ByteCount(offset), + Data: raw[:length], + }) + raw = raw[length:] + default: + t.Fatalf("unexpected frame type 0x%x in chaos-protected payload", raw[0]) + } + } + return crypto, pings, padding +} + +func TestChaosProtectionShreddsCryptoFrame(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient, true) + + const dataLen = 1000 + const paddingLen = 200 + pl := chaosTestPayload(dataLen) + + raw, err := tp.packer.appendChaosProtectedPayload(nil, pl, paddingLen, protocol.Version1) + require.NoError(t, err) + + // The whole point: same size in, same size out. + require.Len(t, raw, int(pl.length+paddingLen)) + + crypto, pings, padding := parseChaosPayload(t, raw) + + // A single frame would mean we did nothing. With no unsplittable frames in + // the way and padding to spare, every attempt lands, so the count is one + // more than the number of attempts. + require.GreaterOrEqual(t, len(crypto), 1+chaosMinAddedCryptoFrames) + require.LessOrEqual(t, len(crypto), 1+chaosMaxAddedCryptoFrames) + require.GreaterOrEqual(t, pings, chaosMinPingFrames) + require.LessOrEqual(t, pings, chaosMaxPingFrames) + require.Positive(t, padding) + + // The CRYPTO fragments must tile the original data exactly once, with no + // gaps and no overlaps, or the peer can't reassemble the ClientHello. + reassembled := make([]byte, dataLen) + covered := make([]bool, dataLen) + for _, cf := range crypto { + for i, b := range cf.Data { + off := int(cf.Offset) + i + require.Less(t, off, dataLen, "fragment runs past the end of the data") + require.False(t, covered[off], "byte %d covered twice", off) + covered[off] = true + reassembled[off] = b + } + } + for i, c := range covered { + require.True(t, c, "byte %d not covered by any fragment", i) + } + for i := range reassembled { + require.Equal(t, byte(i), reassembled[i]) + } +} + +func TestChaosProtectionShufflesOffsets(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient, true) + + // CRYPTO offsets must arrive out of order; always-ascending would itself be + // the fingerprint. + var sawOutOfOrder bool + for range 20 { + raw, err := tp.packer.appendChaosProtectedPayload(nil, chaosTestPayload(1000), 200, protocol.Version1) + require.NoError(t, err) + crypto, _, _ := parseChaosPayload(t, raw) + for i := 1; i < len(crypto); i++ { + if crypto[i].Offset < crypto[i-1].Offset { + sawOutOfOrder = true + } + } + if sawOutOfOrder { + break + } + } + require.True(t, sawOutOfOrder, "CRYPTO frame offsets are never shuffled") +} + +func TestChaosProtectionTinyPaddingBudget(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient, true) + + // With no padding to spend there's no room for extra frame headers, so we + // must fall back to emitting the data unsplit rather than overrun the packet. + for _, paddingLen := range []protocol.ByteCount{0, 1, 2, 5} { + pl := chaosTestPayload(500) + raw, err := tp.packer.appendChaosProtectedPayload(nil, pl, paddingLen, protocol.Version1) + require.NoError(t, err, "padding budget %d", paddingLen) + require.Len(t, raw, int(pl.length+paddingLen), "padding budget %d", paddingLen) + + crypto, _, _ := parseChaosPayload(t, raw) + total := 0 + for _, cf := range crypto { + total += len(cf.Data) + } + require.Equal(t, 500, total, "data lost with padding budget %d", paddingLen) + } +} + +func TestChaosProtectionPreservesOtherFrames(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient, true) + + // Non-CRYPTO frames can't be split, but they must still survive the shuffle. + pl := chaosTestPayload(400) + ping := &wire.PingFrame{} + pl.frames = append(pl.frames, ackhandler.Frame{Frame: ping, Handler: emptyHandler{}}) + pl.length += ping.Length(protocol.Version1) + + raw, err := tp.packer.appendChaosProtectedPayload(nil, pl, 100, protocol.Version1) + require.NoError(t, err) + require.Len(t, raw, int(pl.length+100)) + + crypto, pings, _ := parseChaosPayload(t, raw) + total := 0 + for _, cf := range crypto { + total += len(cf.Data) + } + require.Equal(t, 400, total) + // The carried PING plus the chaos PINGs. + require.GreaterOrEqual(t, pings, chaosMinPingFrames+1) +} + +// available is what's left of a 1250-byte Initial after the long header. +const chaosTestAvailable = 1220 + +func TestChromeCryptoSplitSinglePacket(t *testing.T) { + zero := func(int) int { return 0 } + // Data that fits in one packet is not split at all. + first, last := chromeCryptoSplit(200, 0, chaosTestAvailable, zero) + require.Zero(t, first) + require.Zero(t, last) + + first, last = chromeCryptoSplit(0, 0, chaosTestAvailable, zero) + require.Zero(t, first) + require.Zero(t, last) + + // A packet with no room for two frames can't carry head and tail. + first, last = chromeCryptoSplit(1000, 0, 4, zero) + require.Zero(t, first) + require.Zero(t, last) +} + +func TestChromeCryptoSplitCarriesHeadAndTail(t *testing.T) { + for _, pending := range []protocol.ByteCount{1650, 1700, 1750, 1800, 1850, 2500} { + for r := range chaosFirstFrameLenRandom { + first, last := chromeCryptoSplit(pending, 0, chaosTestAvailable, func(int) int { return r }) + require.Positive(t, first, "pending=%d r=%d", pending, r) + require.Positive(t, last, "pending=%d r=%d", pending, r) + + // The leading frame covers the fixed part of the ClientHello, with + // the randomized offset applied on top. + require.Equal(t, protocol.ByteCount(chaosMinFirstFrameLen+r), first) + + // Head and tail plus their headers have to fit in one packet. + total := first + last + + cryptoFrameHeaderLen(0, first) + + cryptoFrameHeaderLen(pending-last, last) + require.LessOrEqual(t, total, protocol.ByteCount(chaosTestAvailable), + "pending=%d r=%d: first packet overfull", pending, r) + + // Something must be left for the later packets, or nothing is deferred. + require.Positive(t, pending-first-last, "pending=%d r=%d", pending, r) + } + } +} + +func TestChromeCryptoSplitCoversAllData(t *testing.T) { + // Head, tail and the deferred middle must tile the ClientHello exactly: + // a gap stalls the handshake, an overlap is a protocol violation. + for _, pending := range []protocol.ByteCount{1650, 1750, 1850, 2500, 3600} { + first, last := chromeCryptoSplit(pending, 0, chaosTestAvailable, func(n int) int { return n / 2 }) + require.Positive(t, first, "pending=%d", pending) + require.Positive(t, last, "pending=%d", pending) + + covered := make([]bool, pending) + mark := func(off, n protocol.ByteCount) { + for i := off; i < off+n; i++ { + require.False(t, covered[i], "pending=%d: byte %d sent twice", pending, i) + covered[i] = true + } + } + mark(0, first) // head, first packet + mark(pending-last, last) // tail, first packet + mark(first, pending-last-first) // middle, later packets + for i, c := range covered { + require.True(t, c, "pending=%d: byte %d never sent", pending, i) + } + } +} + +func TestChromeCryptoSplitFirstInitialCarriesHeadAndTail(t *testing.T) { + const maxPacketSize protocol.ByteCount = 1250 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient, true) + now := monotime.Now() + + // A ClientHello too large for one Initial, as a post-quantum key share makes it. + const helloLen = 1700 + hello := make([]byte, helloLen) + for i := range hello { + hello[i] = byte(i) + } + tp.initialStream.Write(hello) + + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionInitial). + Return(protocol.PacketNumber(1), protocol.PacketNumberLen1) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, now, gomock.Any()) + // Withholding the Initial ACK is conditional on Handshake keys existing, + // which they do not this early. + tp.sealingManager.EXPECT().GetHandshakeSealer(). + Return(nil, handshake.ErrKeysNotYetAvailable).AnyTimes() + + _, pl := tp.packer.maybeGetCryptoPacket(maxPacketSize, protocol.EncryptionInitial, now, false, false, protocol.Version1) + + var crypto []*wire.CryptoFrame + for _, f := range pl.frames { + if cf, ok := f.Frame.(*wire.CryptoFrame); ok { + crypto = append(crypto, cf) + } + } + require.Len(t, crypto, 2, "the first Initial must carry two CRYPTO frames") + + head, tail := crypto[0], crypto[1] + require.Zero(t, head.Offset, "the first frame must start at the beginning") + require.GreaterOrEqual(t, len(head.Data), chaosMinFirstFrameLen) + require.Less(t, len(head.Data), chaosMinFirstFrameLen+chaosFirstFrameLenRandom) + + // The second frame must reach the very end of the ClientHello, and start + // beyond where the first one stopped: the middle is deferred to a later + // packet, so this one cannot be parsed on its own. + require.Equal(t, protocol.ByteCount(helloLen), tail.Offset+protocol.ByteCount(len(tail.Data))) + require.Greater(t, tail.Offset, head.Offset+protocol.ByteCount(len(head.Data))) +} diff --git a/third_party/quic-go/packet_packer_test.go b/third_party/quic-go/packet_packer_test.go new file mode 100644 index 0000000..3ea8570 --- /dev/null +++ b/third_party/quic-go/packet_packer_test.go @@ -0,0 +1,1086 @@ +package quic + +import ( + "bytes" + "crypto/rand" + "errors" + "testing" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/mocks" + mockackhandler "github.com/apernet/quic-go/internal/mocks/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +const testPackerConnIDLen = 4 + +type testPacketPacker struct { + packer *packetPacker + initialStream *initialCryptoStream + handshakeStream *cryptoStream + datagramQueue *datagramQueue + pnManager *mockackhandler.MockSentPacketHandler + sealingManager *MockSealingManager + framer *MockFrameSource + ackFramer *MockAckFrameSource + retransmissionQueue *retransmissionQueue +} + +// chaosProtection is optional so existing callers stay unchanged; only the +// Chrome chaos protection tests pass it. +func newTestPacketPacker(t *testing.T, mockCtrl *gomock.Controller, pers protocol.Perspective, chaosProtection ...bool) *testPacketPacker { + var chaos bool + if len(chaosProtection) > 0 { + chaos = chaosProtection[0] + } + destConnID := protocol.ParseConnectionID([]byte{1, 2, 3, 4}) + require.Equal(t, testPackerConnIDLen, destConnID.Len()) + initialStream := newInitialCryptoStream(pers == protocol.PerspectiveClient, chaos) + handshakeStream := newCryptoStream() + pnManager := mockackhandler.NewMockSentPacketHandler(mockCtrl) + // Reported for every datagram packed; only the packet number length depends on it. + pnManager.EXPECT().SetLastDatagramPadding(gomock.Any()).AnyTimes() + framer := NewMockFrameSource(mockCtrl) + ackFramer := NewMockAckFrameSource(mockCtrl) + sealingManager := NewMockSealingManager(mockCtrl) + datagramQueue := newDatagramQueue(func() {}, utils.DefaultLogger) + retransmissionQueue := newRetransmissionQueue() + return &testPacketPacker{ + pnManager: pnManager, + initialStream: initialStream, + handshakeStream: handshakeStream, + sealingManager: sealingManager, + framer: framer, + ackFramer: ackFramer, + datagramQueue: datagramQueue, + retransmissionQueue: retransmissionQueue, + packer: newPacketPacker( + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + func() protocol.ConnectionID { return destConnID }, + initialStream, + handshakeStream, + pnManager, + retransmissionQueue, + sealingManager, + framer, + ackFramer, + datagramQueue, + pers, + chaos, + ), + } +} + +// newMockShortHeaderSealer returns a mock short header sealer that seals a short header packet +func newMockShortHeaderSealer(mockCtrl *gomock.Controller) *mocks.MockShortHeaderSealer { + sealer := mocks.NewMockShortHeaderSealer(mockCtrl) + sealer.EXPECT().KeyPhase().Return(protocol.KeyPhaseOne).AnyTimes() + sealer.EXPECT().Overhead().Return(7).AnyTimes() + sealer.EXPECT().EncryptHeader(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + sealer.EXPECT().Seal(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn(func(dst, src []byte, pn protocol.PacketNumber, associatedData []byte) []byte { + return append(src, bytes.Repeat([]byte{'s'}, sealer.Overhead())...) + }).AnyTimes() + return sealer +} + +func parsePacket(t *testing.T, data []byte) (hdrs []*wire.ExtendedHeader, more []byte) { + t.Helper() + for len(data) > 0 { + if !wire.IsLongHeaderPacket(data[0]) { + break + } + hdr, _, more, err := wire.ParsePacket(data) + require.NoError(t, err) + extHdr, err := hdr.ParseExtended(data) + require.NoError(t, err) + require.GreaterOrEqual(t, extHdr.Length+protocol.ByteCount(extHdr.PacketNumberLen), protocol.ByteCount(4)) + data = more + hdrs = append(hdrs, extHdr) + } + return hdrs, data +} + +func parseShortHeaderPacket(t *testing.T, data []byte, connIDLen int) { + t.Helper() + l, _, pnLen, _, err := wire.ParseShortHeader(data, connIDLen) + require.NoError(t, err) + require.GreaterOrEqual(t, len(data)-l+int(pnLen), 4) +} + +func expectAppendFrames(framer *MockFrameSource, controlFrames []ackhandler.Frame, streamFrames []ackhandler.StreamFrame) { + framer.EXPECT().Append(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(cf []ackhandler.Frame, sf []ackhandler.StreamFrame, maxSize protocol.ByteCount, _ monotime.Time, v protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) { + var length protocol.ByteCount + for _, f := range controlFrames { + if length+f.Frame.Length(v) > maxSize { + break + } + length += f.Frame.Length(v) + cf = append(cf, f) + } + for _, f := range streamFrames { + if length+f.Frame.Length(v) > maxSize { + break + } + length += f.Frame.Length(v) + sf = append(sf, f) + } + return cf, sf, length + }, + ) +} + +func generateLargeACKFrame(t *testing.T, minSize protocol.ByteCount) *wire.AckFrame { + t.Helper() + + ack := &wire.AckFrame{ + AckRanges: []wire.AckRange{{Smallest: 1, Largest: 1}}, + DelayTime: 42 * time.Millisecond, + } + var counter int + for ack.Length(protocol.Version1) < minSize { + counter++ + if counter > protocol.MaxNumAckRanges { + t.Fatalf("max number of ACK ranges reached, size: %d", ack.Length(protocol.Version1)) + } + pn := protocol.PacketNumber(1000 * counter) + ack.AckRanges = append([]wire.AckRange{{Smallest: pn, Largest: pn + 100}}, ack.AckRanges...) + } + return ack +} + +func TestPackLongHeaders(t *testing.T) { + skipIfDisableScramblingEnvSet(t) + + t.Run("with Handshake ACK", func(t *testing.T) { + testPackLongHeaders(t, true) + }) + + t.Run("without Handshake ACK", func(t *testing.T) { + testPackLongHeaders(t, false) + }) +} + +func testPackLongHeaders(t *testing.T, includeACK bool) { + const maxPacketSize protocol.ByteCount = 1234 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient) + token := make([]byte, 20) + rand.Read(token) + tp.packer.SetToken(token) + now := monotime.Now() + + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(0x24), protocol.PacketNumberLen3) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(0x24)) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen4) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, now, false) + var numRanges int + if includeACK { + ack := generateLargeACKFrame(t, maxPacketSize-1000) + numRanges = len(ack.AckRanges) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionHandshake, now, false).Return(ack) + } else { + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionHandshake, now, false) + tp.sealingManager.EXPECT().Get0RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + } + clientHello, err := getClientHello("quic-go.net") + require.NoError(t, err) + tp.initialStream.Write(clientHello) + tp.initialStream.Write(make([]byte, 900-len(clientHello))) // add some more data + tp.packer.retransmissionQueue.addHandshake(&wire.PingFrame{}) + + p, err := tp.packer.PackCoalescedPacket(false, maxPacketSize, now, protocol.Version1) + require.NoError(t, err) + require.Equal(t, maxPacketSize, p.buffer.Len()) + require.Len(t, p.longHdrPackets, 2) + require.Nil(t, p.shortHdrPacket) + require.Equal(t, protocol.EncryptionInitial, p.longHdrPackets[0].EncryptionLevel()) + // the ClientHello is split into multiple frames + require.GreaterOrEqual(t, len(p.longHdrPackets[0].frames), 3) + for _, f := range p.longHdrPackets[0].frames { + require.IsType(t, &wire.CryptoFrame{}, f.Frame) + } + require.Equal(t, protocol.EncryptionHandshake, p.longHdrPackets[1].EncryptionLevel()) + require.Len(t, p.longHdrPackets[1].frames, 1) + require.IsType(t, &wire.PingFrame{}, p.longHdrPackets[1].frames[0].Frame) + if includeACK { + require.NotNil(t, p.longHdrPackets[1].ack) + // the ACK frame was truncated + require.Less(t, len(p.longHdrPackets[1].ack.AckRanges), numRanges) + } else { + require.Nil(t, p.longHdrPackets[1].ack) + } + + hdrs, more := parsePacket(t, p.buffer.Data) + require.Len(t, hdrs, 2) + require.Equal(t, protocol.PacketTypeInitial, hdrs[0].Type) + require.Equal(t, token, hdrs[0].Token) + require.Equal(t, protocol.PacketNumber(0x24), hdrs[0].PacketNumber) + require.Equal(t, protocol.PacketNumberLen3, hdrs[0].PacketNumberLen) + require.Equal(t, protocol.PacketTypeHandshake, hdrs[1].Type) + require.Nil(t, hdrs[1].Token) + require.Equal(t, protocol.PacketNumber(0x42), hdrs[1].PacketNumber) + require.Equal(t, protocol.PacketNumberLen4, hdrs[1].PacketNumberLen) + require.Empty(t, more) +} + +func TestPackCoalescedAckOnlyPacketNothingToSend(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + // the packet number is not popped + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, gomock.Any(), true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionHandshake, gomock.Any(), true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true) + p, err := tp.packer.PackCoalescedPacket(true, 1234, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Nil(t, p) +} + +func TestPackInitialAckOnlyPacket(t *testing.T) { + t.Run("client", func(t *testing.T) { testPackInitialAckOnlyPacket(t, protocol.PerspectiveClient) }) + t.Run("server", func(t *testing.T) { testPackInitialAckOnlyPacket(t, protocol.PerspectiveServer) }) +} + +func testPackInitialAckOnlyPacket(t *testing.T, pers protocol.Perspective) { + const maxPacketSize protocol.ByteCount = 1234 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, pers) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + ack := &wire.AckFrame{AckRanges: []wire.AckRange{{Smallest: 1, Largest: 10}}} + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, gomock.Any(), true).Return(ack) + p, err := tp.packer.PackCoalescedPacket(true, maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p) + require.Len(t, p.longHdrPackets, 1) + require.Equal(t, protocol.EncryptionInitial, p.longHdrPackets[0].EncryptionLevel()) + require.Equal(t, ack, p.longHdrPackets[0].ack) + require.Empty(t, p.longHdrPackets[0].frames) + // only the client needs to pad Initial packets + switch pers { + case protocol.PerspectiveClient: + require.Equal(t, maxPacketSize, p.buffer.Len()) + case protocol.PerspectiveServer: + require.Less(t, p.buffer.Len(), protocol.ByteCount(100)) + } + hdrs, more := parsePacket(t, p.buffer.Data) + require.Empty(t, more) + require.Len(t, hdrs, 1) + require.Equal(t, protocol.PacketTypeInitial, hdrs[0].Type) +} + +func TestPack1RTTAckOnlyPacket(t *testing.T) { + const maxPacketSize protocol.ByteCount = 1300 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + ack := &wire.AckFrame{AckRanges: []wire.AckRange{{Smallest: 1, Largest: 10}}} + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true).Return(ack) + p, buffer, err := tp.packer.PackAckOnlyPacket(maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Equal(t, ack, p.Ack) + require.Empty(t, p.Frames) + parsePacket(t, buffer.Data) +} + +func TestPack0RTTPacket(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient) + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().Get0RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, gomock.Any(), true) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption0RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption0RTT).Return(protocol.PacketNumber(0x42)) + cf := ackhandler.Frame{Frame: &wire.MaxDataFrame{MaximumData: 0x1337}} + tp.framer.EXPECT().HasData().Return(true) + // TODO: check sizes + tp.framer.EXPECT().Append(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(fs []ackhandler.Frame, sf []ackhandler.StreamFrame, _ protocol.ByteCount, _ monotime.Time, _ protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) { + return append(fs, cf), sf, cf.Frame.Length(protocol.Version1) + }, + ) + p, err := tp.packer.PackCoalescedPacket(false, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p) + require.Len(t, p.longHdrPackets, 1) + require.Equal(t, protocol.PacketType0RTT, p.longHdrPackets[0].header.Type) + require.Equal(t, protocol.Encryption0RTT, p.longHdrPackets[0].EncryptionLevel()) + require.Len(t, p.longHdrPackets[0].frames, 1) + require.Equal(t, cf.Frame, p.longHdrPackets[0].frames[0].Frame) + require.NotNil(t, p.longHdrPackets[0].frames[0].Handler) +} + +// ACK frames can't be sent in 0-RTT packets +func TestPack0RTTPacketNoACK(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient) + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, gomock.Any(), true) + // no further calls to get an ACK frame + p, err := tp.packer.PackCoalescedPacket(true, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Nil(t, p) +} + +func TestPackCoalescedAppData(t *testing.T) { + t.Run("with large ACK", func(t *testing.T) { + testPackCoalescedAppData(t, true) + }) + + t.Run("without ACK", func(t *testing.T) { + testPackCoalescedAppData(t, false) + }) +} + +func testPackCoalescedAppData(t *testing.T, withAck bool) { + const maxPacketSize protocol.ByteCount = 1234 + + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x24), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x24)) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(nil, handshake.ErrKeysDropped) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionHandshake, gomock.Any(), false) + + var numRanges int + if withAck { + // The ACK is too large and needs to be truncated + ack := generateLargeACKFrame(t, maxPacketSize-1000) + numRanges = len(ack.AckRanges) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false).Return(ack) + } else { + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false) + } + + handshakeData := make([]byte, 1000) + rand.Read(handshakeData) + tp.handshakeStream.Write(handshakeData) + expectAppendFrames(tp.framer, nil, []ackhandler.StreamFrame{{Frame: &wire.StreamFrame{Data: []byte("foobar")}}}) + + p, err := tp.packer.PackCoalescedPacket(false, maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Len(t, p.longHdrPackets, 1) + require.Equal(t, protocol.EncryptionHandshake, p.longHdrPackets[0].EncryptionLevel()) + require.Len(t, p.longHdrPackets[0].frames, 1) + require.Equal(t, handshakeData, p.longHdrPackets[0].frames[0].Frame.(*wire.CryptoFrame).Data) + require.NotNil(t, p.shortHdrPacket) + require.Empty(t, p.shortHdrPacket.Frames) + if withAck { + require.NotNil(t, p.shortHdrPacket.Ack) + require.Less(t, len(p.shortHdrPacket.Ack.AckRanges), numRanges) + require.LessOrEqual(t, len(p.buffer.Data), int(maxPacketSize)) + require.Empty(t, p.shortHdrPacket.StreamFrames) + } else { + require.Nil(t, p.shortHdrPacket.Ack) + require.Less(t, len(p.buffer.Data), int(maxPacketSize)) + require.Len(t, p.shortHdrPacket.StreamFrames, 1) + require.Equal(t, []byte("foobar"), p.shortHdrPacket.StreamFrames[0].Frame.Data) + } + + hdrs, more := parsePacket(t, p.buffer.Data) + require.Len(t, hdrs, 1) + require.Equal(t, protocol.PacketTypeHandshake, hdrs[0].Type) + require.NotEmpty(t, more) + parseShortHeaderPacket(t, more, testPackerConnIDLen) +} + +func TestPackConnectionCloseCoalesced(t *testing.T) { + t.Run("client", func(t *testing.T) { testPackConnectionCloseCoalesced(t, protocol.PerspectiveClient) }) + t.Run("server", func(t *testing.T) { testPackConnectionCloseCoalesced(t, protocol.PerspectiveServer) }) +} + +func testPackConnectionCloseCoalesced(t *testing.T, pers protocol.Perspective) { + const maxPacketSize protocol.ByteCount = 1234 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, pers) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(1), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(1)) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(2), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(2)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + switch pers { + case protocol.PerspectiveClient: + tp.sealingManager.EXPECT().Get0RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption0RTT).Return(protocol.PacketNumber(3), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption0RTT).Return(protocol.PacketNumber(3)) + case protocol.PerspectiveServer: + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(3), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(3)) + } + p, err := tp.packer.PackApplicationClose(&qerr.ApplicationError{ + ErrorCode: 0x1337, + ErrorMessage: "foobar", + }, maxPacketSize, protocol.Version1) + require.NoError(t, err) + switch pers { + case protocol.PerspectiveClient: + require.Len(t, p.longHdrPackets, 3) + require.Nil(t, p.shortHdrPacket) + case protocol.PerspectiveServer: + require.Len(t, p.longHdrPackets, 2) + require.NotNil(t, p.shortHdrPacket) + } + // for Initial packets, the error code is replace with a transport error of type APPLICATION_ERROR + require.Equal(t, protocol.PacketTypeInitial, p.longHdrPackets[0].header.Type) + require.Equal(t, protocol.PacketNumber(1), p.longHdrPackets[0].header.PacketNumber) + require.Len(t, p.longHdrPackets[0].frames, 1) + require.IsType(t, &wire.ConnectionCloseFrame{}, p.longHdrPackets[0].frames[0].Frame) + ccf := p.longHdrPackets[0].frames[0].Frame.(*wire.ConnectionCloseFrame) + require.False(t, ccf.IsApplicationError) + require.Equal(t, uint64(qerr.ApplicationErrorErrorCode), ccf.ErrorCode) + require.Empty(t, ccf.ReasonPhrase) + // for Handshake packets, the error code is replace with a transport error of type APPLICATION_ERROR + require.Equal(t, protocol.PacketTypeHandshake, p.longHdrPackets[1].header.Type) + require.Equal(t, protocol.PacketNumber(2), p.longHdrPackets[1].header.PacketNumber) + require.Len(t, p.longHdrPackets[1].frames, 1) + require.IsType(t, &wire.ConnectionCloseFrame{}, p.longHdrPackets[1].frames[0].Frame) + ccf = p.longHdrPackets[1].frames[0].Frame.(*wire.ConnectionCloseFrame) + require.False(t, ccf.IsApplicationError) + require.Equal(t, uint64(qerr.ApplicationErrorErrorCode), ccf.ErrorCode) + require.Empty(t, ccf.ReasonPhrase) + + // for application-data packet number space (1-RTT for the server, 0-RTT for the client), + // the application-level error code is sent + + switch pers { + case protocol.PerspectiveClient: + require.Equal(t, protocol.PacketNumber(3), p.longHdrPackets[2].header.PacketNumber) + require.Len(t, p.longHdrPackets[2].frames, 1) + require.IsType(t, &wire.ConnectionCloseFrame{}, p.longHdrPackets[2].frames[0].Frame) + ccf = p.longHdrPackets[2].frames[0].Frame.(*wire.ConnectionCloseFrame) + case protocol.PerspectiveServer: + require.Equal(t, protocol.PacketNumber(3), p.shortHdrPacket.PacketNumber) + require.Len(t, p.shortHdrPacket.Frames, 1) + require.IsType(t, &wire.ConnectionCloseFrame{}, p.shortHdrPacket.Frames[0].Frame) + ccf = p.shortHdrPacket.Frames[0].Frame.(*wire.ConnectionCloseFrame) + } + require.True(t, ccf.IsApplicationError) + require.Equal(t, uint64(0x1337), ccf.ErrorCode) + require.Equal(t, "foobar", ccf.ReasonPhrase) + + // the client needs to pad this packet to the max packet size + switch pers { + case protocol.PerspectiveClient: + require.Equal(t, maxPacketSize, p.buffer.Len()) + case protocol.PerspectiveServer: + require.Less(t, p.buffer.Len(), protocol.ByteCount(100)) + } +} + +func TestPackConnectionCloseCryptoError(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(nil, handshake.ErrKeysDropped) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + quicErr := qerr.NewLocalCryptoError(0x42, errors.New("crypto error")) + quicErr.FrameType = 0x1234 + p, err := tp.packer.PackConnectionClose(quicErr, protocol.MaxByteCount, protocol.Version1) + require.NoError(t, err) + require.Len(t, p.longHdrPackets, 1) + require.Equal(t, protocol.PacketTypeHandshake, p.longHdrPackets[0].header.Type) + require.Len(t, p.longHdrPackets[0].frames, 1) + require.IsType(t, &wire.ConnectionCloseFrame{}, p.longHdrPackets[0].frames[0].Frame) + ccf := p.longHdrPackets[0].frames[0].Frame.(*wire.ConnectionCloseFrame) + require.False(t, ccf.IsApplicationError) + require.Equal(t, uint64(0x100+0x42), ccf.ErrorCode) + require.Equal(t, uint64(0x1234), ccf.FrameType) + // for crypto errors, the reason phrase is cleared + require.Empty(t, ccf.ReasonPhrase) +} + +func TestPackConnectionClose1RTT(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(nil, handshake.ErrKeysDropped) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(nil, handshake.ErrKeysDropped) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + // expect no framer.PopStreamFrames + p, err := tp.packer.PackConnectionClose(&qerr.TransportError{ + ErrorCode: qerr.CryptoBufferExceeded, + ErrorMessage: "foo", + }, protocol.MaxByteCount, protocol.Version1) + require.NoError(t, err) + require.Empty(t, p.longHdrPackets) + require.Len(t, p.shortHdrPacket.Frames, 1) + require.IsType(t, &wire.ConnectionCloseFrame{}, p.shortHdrPacket.Frames[0].Frame) + ccf := p.shortHdrPacket.Frames[0].Frame.(*wire.ConnectionCloseFrame) + require.False(t, ccf.IsApplicationError) + require.Equal(t, uint64(qerr.CryptoBufferExceeded), ccf.ErrorCode) + require.Equal(t, "foo", ccf.ReasonPhrase) +} + +func TestPack1RTTPacketNothingToSend(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + // don't expect any calls to PopPacketNumber + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true) + tp.framer.EXPECT().HasData() + _, err := tp.packer.AppendPacket(getPacketBuffer(), protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.ErrorIs(t, err, errNothingToPack) +} + +func TestPack1RTTPacketWithData(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false) + f := &wire.StreamFrame{ + StreamID: 5, + Data: []byte{0xde, 0xca, 0xfb, 0xad}, + } + expectAppendFrames( + tp.framer, + []ackhandler.Frame{ + {Frame: &wire.ResetStreamFrame{}, Handler: &mtuFinderAckHandler{}}, // set any non-nil ackhandler.FrameHandler + {Frame: &wire.MaxDataFrame{}}, + }, + []ackhandler.StreamFrame{{Frame: f}}, + ) + buffer := getPacketBuffer() + buffer.Data = append(buffer.Data, []byte("foobar")...) + p, err := tp.packer.AppendPacket(buffer, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + b, err := f.Append(nil, protocol.Version1) + require.NoError(t, err) + require.Len(t, p.StreamFrames, 1) + var sawResetStream, sawMaxData bool + for _, frame := range p.Frames { + switch frame.Frame.(type) { + case *wire.ResetStreamFrame: + sawResetStream = true + require.Equal(t, frame.Handler, &mtuFinderAckHandler{}) + case *wire.MaxDataFrame: + sawMaxData = true + require.NotNil(t, frame.Handler) + require.NotEqual(t, frame.Handler, &mtuFinderAckHandler{}) + } + } + require.True(t, sawResetStream) + require.True(t, sawMaxData) + require.Equal(t, f.StreamID, p.StreamFrames[0].Frame.StreamID) + require.Equal(t, buffer.Data[:6], []byte("foobar")) // make sure the packet was actually appended + require.Contains(t, string(buffer.Data), string(b)) +} + +func TestPack1RTTPacketWithACK(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + ack := &wire.AckFrame{AckRanges: []wire.AckRange{{Largest: 42, Smallest: 1}}} + tp.framer.EXPECT().HasData() + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true).Return(ack) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + p, err := tp.packer.AppendPacket(getPacketBuffer(), protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Equal(t, ack, p.Ack) +} + +func TestPackPathChallengeAndPathResponse(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false) + frames := []ackhandler.Frame{ + {Frame: &wire.PathChallengeFrame{}}, + {Frame: &wire.PathResponseFrame{}}, + {Frame: &wire.DataBlockedFrame{}}, + } + expectAppendFrames(tp.framer, frames, nil) + buffer := getPacketBuffer() + p, err := tp.packer.AppendPacket(buffer, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Len(t, p.Frames, 3) + var sawPathChallenge, sawPathResponse bool + for _, f := range p.Frames { + switch f.Frame.(type) { + case *wire.PathChallengeFrame: + sawPathChallenge = true + // this means that the frame won't be retransmitted. + require.Nil(t, f.Handler) + case *wire.PathResponseFrame: + sawPathResponse = true + // this means that the frame won't be retransmitted. + require.Nil(t, f.Handler) + default: + require.NotNil(t, f.Handler) + } + } + require.True(t, sawPathChallenge) + require.True(t, sawPathResponse) + require.NotZero(t, buffer.Len()) +} + +func TestPackDatagramFrames(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.datagramQueue.Add(&wire.DatagramFrame{ + DataLenPresent: true, + Data: []byte("foobar"), + }) + tp.framer.EXPECT().HasData() + buffer := getPacketBuffer() + p, err := tp.packer.AppendPacket(buffer, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Len(t, p.Frames, 1) + require.IsType(t, &wire.DatagramFrame{}, p.Frames[0].Frame) + require.Equal(t, []byte("foobar"), p.Frames[0].Frame.(*wire.DatagramFrame).Data) + require.NotEmpty(t, buffer.Data) +} + +func TestPackLargeDatagramFrame(t *testing.T) { + // If a packet contains an ACK, and doesn't have enough space for the DATAGRAM frame, + // it should be skipped. It will be packed in the next packet. + const maxPacketSize = 1000 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true).Return(&wire.AckFrame{AckRanges: []wire.AckRange{{Largest: 100}}}) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + f := &wire.DatagramFrame{DataLenPresent: true, Data: make([]byte, maxPacketSize-10)} + tp.datagramQueue.Add(f) + tp.framer.EXPECT().HasData() + buffer := getPacketBuffer() + p, err := tp.packer.AppendPacket(buffer, maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p.Ack) + require.Empty(t, p.Frames) + require.NotEmpty(t, buffer.Data) + require.Equal(t, f, tp.datagramQueue.Peek()) // make sure the frame is still there + + // Now try packing again, but with a smaller packet size. + // The DATAGRAM frame should now be dropped, as we can't expect to ever be able tosend it out. + const newMaxPacketSize = maxPacketSize - 10 + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), true) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x43), protocol.PacketNumberLen2) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData() + buffer = getPacketBuffer() + p, err = tp.packer.AppendPacket(buffer, newMaxPacketSize, monotime.Now(), protocol.Version1) + require.ErrorIs(t, err, errNothingToPack) + require.Nil(t, tp.datagramQueue.Peek()) // make sure the frame is gone +} + +func TestPackRetransmissions(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + f := &wire.CryptoFrame{Data: []byte("Initial")} + tp.retransmissionQueue.addInitial(f) + tp.retransmissionQueue.addHandshake(&wire.CryptoFrame{Data: []byte("Handshake")}) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionInitial).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionInitial, gomock.Any(), false) + p, err := tp.packer.PackCoalescedPacket(false, 1000, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Len(t, p.longHdrPackets, 1) + require.Equal(t, protocol.EncryptionInitial, p.longHdrPackets[0].EncryptionLevel()) + require.Len(t, p.longHdrPackets[0].frames, 1) + require.Equal(t, f, p.longHdrPackets[0].frames[0].Frame) + require.NotNil(t, p.longHdrPackets[0].frames[0].Handler) +} + +func packMaxNumNonAckElicitingAcks(t *testing.T, tp *testPacketPacker, mockCtrl *gomock.Controller, maxPacketSize protocol.ByteCount) { + t.Helper() + for range protocol.MaxNonAckElicitingAcks { + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false).Return( + &wire.AckFrame{AckRanges: []wire.AckRange{{Smallest: 1, Largest: 1}}}, + ) + expectAppendFrames(tp.framer, nil, nil) + p, err := tp.packer.AppendPacket(getPacketBuffer(), maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p.Ack) + require.Empty(t, p.Frames) + } +} + +func TestPackEvery20thPacketAckEliciting(t *testing.T) { + const maxPacketSize = 1000 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + + // send the maximum number of non-ACK-eliciting packets + packMaxNumNonAckElicitingAcks(t, tp, mockCtrl, maxPacketSize) + + // Now there's nothing to send, so we shouldn't generate a packet just to send a PING + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false) + tp.framer.EXPECT().HasData().Return(true) + expectAppendFrames(tp.framer, nil, nil) + _, err := tp.packer.AppendPacket(getPacketBuffer(), maxPacketSize, monotime.Now(), protocol.Version1) + require.ErrorIs(t, err, errNothingToPack) + + // Now we have an ACK to send. We should bundle a PING to make the packet ack-eliciting. + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false).Return( + &wire.AckFrame{AckRanges: []wire.AckRange{{Smallest: 1, Largest: 1}}}, + ) + expectAppendFrames(tp.framer, nil, nil) + p, err := tp.packer.AppendPacket(getPacketBuffer(), maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Len(t, p.Frames, 1) + require.Equal(t, &wire.PingFrame{}, p.Frames[0].Frame) + require.Nil(t, p.Frames[0].Handler) // make sure the PING is not retransmitted if lost + + // make sure the next packet doesn't contain another PING + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false).Return( + &wire.AckFrame{AckRanges: []wire.AckRange{{Smallest: 1, Largest: 1}}}, + ) + expectAppendFrames(tp.framer, nil, nil) + p, err = tp.packer.AppendPacket(getPacketBuffer(), maxPacketSize, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p.Ack) + require.Empty(t, p.Frames) +} + +func TestPackLongHeaderPadToAtLeast4Bytes(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen1) + tp.pnManager.EXPECT().PopPacketNumber(protocol.EncryptionHandshake).Return(protocol.PacketNumber(0x42)) + + sealer := newMockShortHeaderSealer(mockCtrl) + tp.sealingManager.EXPECT().GetInitialSealer().Return(nil, handshake.ErrKeysDropped) + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(sealer, nil) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(nil, handshake.ErrKeysNotYetAvailable) + tp.retransmissionQueue.addHandshake(&wire.PingFrame{}) + tp.ackFramer.EXPECT().GetAckFrame(protocol.EncryptionHandshake, gomock.Any(), false) + + packet, err := tp.packer.PackCoalescedPacket(false, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, packet) + require.Len(t, packet.longHdrPackets, 1) + require.Nil(t, packet.shortHdrPacket) + + hdr, _, _, err := wire.ParsePacket(packet.buffer.Data) + require.NoError(t, err) + data := packet.buffer.Data + extHdr, err := hdr.ParseExtended(data) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumberLen1, extHdr.PacketNumberLen) + + data = data[extHdr.ParsedLen():] + require.Len(t, data, 4-1 /* packet number length */ +sealer.Overhead()) + // first bytes should be 2 PADDING frames... + require.Equal(t, []byte{0, 0}, data[:2]) + // ...followed by the PING frame + frameParser := wire.NewFrameParser(false, false, false) + + frameType, lt, err := frameParser.ParseType(data[2:], protocol.EncryptionHandshake) + require.NoError(t, err) + require.Equal(t, 1, lt) + frame, l, err := frameParser.ParseLessCommonFrame(frameType, data[2+lt:], protocol.Version1) + require.NoError(t, err) + require.IsType(t, &wire.PingFrame{}, frame) + require.Zero(t, l) + require.Equal(t, sealer.Overhead(), len(data)-2-lt) +} + +func TestPackShortHeaderPadToAtLeast4Bytes(t *testing.T) { + // small stream ID, such that only a single byte is consumed + f := &wire.StreamFrame{StreamID: 0x10, Fin: true} + require.Equal(t, protocol.ByteCount(2), f.Length(protocol.Version1)) + + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen1) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + sealer := newMockShortHeaderSealer(mockCtrl) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(sealer, nil) + tp.framer.EXPECT().HasData().Return(true) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false) + expectAppendFrames(tp.framer, nil, []ackhandler.StreamFrame{{Frame: f}}) + + buffer := getPacketBuffer() + _, err := tp.packer.AppendPacket(buffer, protocol.MaxByteCount, monotime.Now(), protocol.Version1) + require.NoError(t, err) + // cut off the tag that the mock sealer added + buffer.Data = buffer.Data[:buffer.Len()-protocol.ByteCount(sealer.Overhead())] + data := buffer.Data + + l, _, pnLen, _, err := wire.ParseShortHeader(data, testPackerConnIDLen) + require.NoError(t, err) + payload := data[l:] + require.Equal(t, protocol.PacketNumberLen1, pnLen) + require.Equal(t, 4-1 /* packet number length */, len(payload)) + // the first byte of the payload should be a PADDING frame... + require.Equal(t, byte(0), payload[0]) + + // ... followed by the STREAM frame + frameParser := wire.NewFrameParser(false, false, false) + frameType, l, err := frameParser.ParseType(payload[1:], protocol.Encryption1RTT) + require.NoError(t, err) + require.Equal(t, 1, l) + require.True(t, frameType.IsStreamFrameType()) + + frame, frameLen, err := wire.ParseStreamFrame(payload[1+l:], frameType, protocol.Version1) + require.NoError(t, err) + require.Equal(t, f, frame) + require.Equal(t, len(payload)-2, frameLen) +} + +func TestPackInitialProbePacket(t *testing.T) { + t.Run("client", func(t *testing.T) { + t.Setenv(disableClientHelloScramblingEnv, "true") + testPackProbePacket(t, protocol.EncryptionInitial, protocol.PerspectiveClient) + }) + t.Run("server", func(t *testing.T) { + testPackProbePacket(t, protocol.EncryptionInitial, protocol.PerspectiveServer) + }) +} + +func TestPackHandshakeProbePacket(t *testing.T) { + t.Run("client", func(t *testing.T) { + testPackProbePacket(t, protocol.EncryptionHandshake, protocol.PerspectiveClient) + }) + t.Run("server", func(t *testing.T) { + testPackProbePacket(t, protocol.EncryptionHandshake, protocol.PerspectiveServer) + }) +} + +func testPackProbePacket(t *testing.T, encLevel protocol.EncryptionLevel, perspective protocol.Perspective) { + const maxPacketSize protocol.ByteCount = 1234 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, perspective) + + var cryptoData []byte + switch encLevel { + case protocol.EncryptionInitial: + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + var err error + cryptoData, err = getClientHello("") + require.NoError(t, err) + tp.packer.initialStream.Write(cryptoData) + case protocol.EncryptionHandshake: + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + cryptoData = []byte("foobar") + tp.packer.handshakeStream.Write(cryptoData) + } + tp.ackFramer.EXPECT().GetAckFrame(encLevel, gomock.Any(), false) + tp.pnManager.EXPECT().PeekPacketNumber(encLevel).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(encLevel).Return(protocol.PacketNumber(0x42)) + + p, err := tp.packer.PackPTOProbePacket(encLevel, maxPacketSize, false, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p) + require.Len(t, p.longHdrPackets, 1) + packet := p.longHdrPackets[0] + require.Equal(t, encLevel, packet.EncryptionLevel()) + if encLevel == protocol.EncryptionInitial { + require.GreaterOrEqual(t, p.buffer.Len(), protocol.ByteCount(protocol.MinInitialPacketSize)) + require.Equal(t, maxPacketSize, p.buffer.Len()) + } + require.Len(t, packet.frames, 1) + require.Equal(t, cryptoData, packet.frames[0].Frame.(*wire.CryptoFrame).Data) + hdrs, more := parsePacket(t, p.buffer.Data) + require.Len(t, hdrs, 1) + switch encLevel { + case protocol.EncryptionInitial: + require.Equal(t, protocol.PacketTypeInitial, hdrs[0].Type) + case protocol.EncryptionHandshake: + require.Equal(t, protocol.PacketTypeHandshake, hdrs[0].Type) + } + require.Empty(t, more) +} + +func TestPack1RTTProbePacket(t *testing.T) { + const maxPacketSize protocol.ByteCount = 999 + + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.ackFramer.EXPECT().GetAckFrame(protocol.Encryption1RTT, gomock.Any(), false) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x42)) + tp.framer.EXPECT().HasData().Return(true) + tp.framer.EXPECT().Append(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), protocol.Version1).DoAndReturn( + func(cf []ackhandler.Frame, sf []ackhandler.StreamFrame, size protocol.ByteCount, _ monotime.Time, v protocol.Version) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) { + f, split := (&wire.StreamFrame{Data: make([]byte, 2*maxPacketSize)}).MaybeSplitOffFrame(size, v) + require.True(t, split) + return cf, append(sf, ackhandler.StreamFrame{Frame: f}), f.Length(v) + }, + ) + + p, err := tp.packer.PackPTOProbePacket(protocol.Encryption1RTT, maxPacketSize, false, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, p) + require.True(t, p.IsOnlyShortHeaderPacket()) + require.Empty(t, p.longHdrPackets) + require.NotNil(t, p.shortHdrPacket) + packet := p.shortHdrPacket + require.Empty(t, packet.Frames) + require.Len(t, packet.StreamFrames, 1) + require.Equal(t, maxPacketSize, packet.Length) +} + +func TestPackPTOProbePacketNothingToPack(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testPackPTOProbePacketNothingToPack(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testPackPTOProbePacketNothingToPack(t, protocol.EncryptionHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testPackPTOProbePacketNothingToPack(t, protocol.Encryption1RTT) + }) +} + +func testPackPTOProbePacketNothingToPack(t *testing.T, encLevel protocol.EncryptionLevel) { + const maxPacketSize protocol.ByteCount = 1234 + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + + switch encLevel { + case protocol.EncryptionInitial: + tp.sealingManager.EXPECT().GetInitialSealer().Return(newMockShortHeaderSealer(mockCtrl), nil).Times(2) + case protocol.EncryptionHandshake: + tp.sealingManager.EXPECT().GetHandshakeSealer().Return(newMockShortHeaderSealer(mockCtrl), nil).Times(2) + case protocol.Encryption1RTT: + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil).Times(2) + tp.framer.EXPECT().HasData().Times(2) + } + tp.pnManager.EXPECT().PeekPacketNumber(encLevel).Return(protocol.PacketNumber(0x42), protocol.PacketNumberLen2).MaxTimes(2) + tp.ackFramer.EXPECT().GetAckFrame(encLevel, gomock.Any(), true).Times(2) + + // don't force a PING to be sent + packet, err := tp.packer.PackPTOProbePacket(encLevel, maxPacketSize, false, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.Nil(t, packet) + + // now force a PING to be sent + tp.pnManager.EXPECT().PopPacketNumber(encLevel).Return(protocol.PacketNumber(0x42)) + packet, err = tp.packer.PackPTOProbePacket(encLevel, maxPacketSize, true, monotime.Now(), protocol.Version1) + require.NoError(t, err) + require.NotNil(t, packet) + var frames []ackhandler.Frame + switch encLevel { + case protocol.EncryptionInitial, protocol.EncryptionHandshake: + require.Len(t, packet.longHdrPackets, 1) + require.Nil(t, packet.shortHdrPacket) + require.Equal(t, encLevel, packet.longHdrPackets[0].EncryptionLevel()) + frames = packet.longHdrPackets[0].frames + case protocol.Encryption1RTT: + require.Empty(t, packet.longHdrPackets) + require.NotNil(t, packet.shortHdrPacket) + frames = packet.shortHdrPacket.Frames + } + + require.Len(t, frames, 1) + require.Equal(t, &wire.PingFrame{}, frames[0].Frame) + require.Equal(t, emptyHandler{}, frames[0].Handler) +} + +func TestPackMTUProbePacket(t *testing.T) { + const ( + maxPacketSize protocol.ByteCount = 1000 + probePacketSize = maxPacketSize + 42 + ) + + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveClient) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x43), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x43)) + ping := ackhandler.Frame{Frame: &wire.PingFrame{}} + p, buffer, err := tp.packer.PackMTUProbePacket(ping, probePacketSize, protocol.Version1) + require.NoError(t, err) + require.Equal(t, probePacketSize, p.Length) + require.Equal(t, protocol.PacketNumber(0x43), p.PacketNumber) + require.Len(t, buffer.Data, int(probePacketSize)) + require.True(t, p.IsPathMTUProbePacket) + require.False(t, p.IsPathProbePacket) +} + +func TestPackPathProbePacket(t *testing.T) { + mockCtrl := gomock.NewController(t) + tp := newTestPacketPacker(t, mockCtrl, protocol.PerspectiveServer) + tp.sealingManager.EXPECT().Get1RTTSealer().Return(newMockShortHeaderSealer(mockCtrl), nil) + tp.pnManager.EXPECT().PeekPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x43), protocol.PacketNumberLen2) + tp.pnManager.EXPECT().PopPacketNumber(protocol.Encryption1RTT).Return(protocol.PacketNumber(0x43)) + + p, buf, err := tp.packer.PackPathProbePacket( + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + []ackhandler.Frame{ + {Frame: &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}}, + {Frame: &wire.PathResponseFrame{Data: [8]byte{8, 7, 6, 5, 4, 3, 2, 1}}}, + }, + protocol.Version1, + ) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumber(0x43), p.PacketNumber) + require.Nil(t, p.Ack) + require.Empty(t, p.StreamFrames) + require.Len(t, p.Frames, 2) + // the frame order is randomized + frames := []wire.Frame{p.Frames[0].Frame, p.Frames[1].Frame} + require.Contains(t, frames, &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}) + require.Contains(t, frames, &wire.PathResponseFrame{Data: [8]byte{8, 7, 6, 5, 4, 3, 2, 1}}) + require.Len(t, buf.Data, protocol.MinInitialPacketSize) + require.True(t, p.IsPathProbePacket) + require.False(t, p.IsPathMTUProbePacket) +} diff --git a/third_party/quic-go/packet_unpacker.go b/third_party/quic-go/packet_unpacker.go new file mode 100644 index 0000000..6ec850b --- /dev/null +++ b/third_party/quic-go/packet_unpacker.go @@ -0,0 +1,222 @@ +package quic + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +type headerDecryptor interface { + DecryptHeader(sample []byte, firstByte *byte, pnBytes []byte) +} + +type headerParseError struct { + err error +} + +func (e *headerParseError) Unwrap() error { + return e.err +} + +func (e *headerParseError) Error() string { + return e.err.Error() +} + +type unpackedPacket struct { + hdr *wire.ExtendedHeader + encryptionLevel protocol.EncryptionLevel + data []byte +} + +// The packetUnpacker unpacks QUIC packets. +type packetUnpacker struct { + cs handshake.CryptoSetup + + shortHdrConnIDLen int +} + +var _ unpacker = &packetUnpacker{} + +func newPacketUnpacker(cs handshake.CryptoSetup, shortHdrConnIDLen int) *packetUnpacker { + return &packetUnpacker{ + cs: cs, + shortHdrConnIDLen: shortHdrConnIDLen, + } +} + +// UnpackLongHeader unpacks a Long Header packet. +// If the reserved bits are invalid, the error is wire.ErrInvalidReservedBits. +// If any other error occurred when parsing the header, the error is of type headerParseError. +// If decrypting the payload fails for any reason, the error is the error returned by the AEAD. +func (u *packetUnpacker) UnpackLongHeader(hdr *wire.Header, data []byte) (*unpackedPacket, error) { + var encLevel protocol.EncryptionLevel + var extHdr *wire.ExtendedHeader + var decrypted []byte + //nolint:exhaustive // Retry packets can't be unpacked. + switch hdr.Type { + case protocol.PacketTypeInitial: + encLevel = protocol.EncryptionInitial + opener, err := u.cs.GetInitialOpener() + if err != nil { + return nil, err + } + extHdr, decrypted, err = u.unpackLongHeaderPacket(opener, hdr, data) + if err != nil { + return nil, err + } + case protocol.PacketTypeHandshake: + encLevel = protocol.EncryptionHandshake + opener, err := u.cs.GetHandshakeOpener() + if err != nil { + return nil, err + } + extHdr, decrypted, err = u.unpackLongHeaderPacket(opener, hdr, data) + if err != nil { + return nil, err + } + case protocol.PacketType0RTT: + encLevel = protocol.Encryption0RTT + opener, err := u.cs.Get0RTTOpener() + if err != nil { + return nil, err + } + extHdr, decrypted, err = u.unpackLongHeaderPacket(opener, hdr, data) + if err != nil { + return nil, err + } + default: + return nil, fmt.Errorf("unknown packet type: %s", hdr.Type) + } + + if len(decrypted) == 0 { + return nil, &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "empty packet", + } + } + + return &unpackedPacket{ + hdr: extHdr, + encryptionLevel: encLevel, + data: decrypted, + }, nil +} + +func (u *packetUnpacker) UnpackShortHeader(rcvTime monotime.Time, data []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error) { + opener, err := u.cs.Get1RTTOpener() + if err != nil { + return 0, 0, 0, nil, err + } + pn, pnLen, kp, decrypted, err := u.unpackShortHeaderPacket(opener, rcvTime, data) + if err != nil { + return 0, 0, 0, nil, err + } + if len(decrypted) == 0 { + return 0, 0, 0, nil, &qerr.TransportError{ + ErrorCode: qerr.ProtocolViolation, + ErrorMessage: "empty packet", + } + } + return pn, pnLen, kp, decrypted, nil +} + +func (u *packetUnpacker) unpackLongHeaderPacket(opener handshake.LongHeaderOpener, hdr *wire.Header, data []byte) (*wire.ExtendedHeader, []byte, error) { + extHdr, parseErr := u.unpackLongHeader(opener, hdr, data) + // If the reserved bits are set incorrectly, we still need to continue unpacking. + // This avoids a timing side-channel, which otherwise might allow an attacker + // to gain information about the header encryption. + if parseErr != nil && parseErr != wire.ErrInvalidReservedBits { + return nil, nil, parseErr + } + extHdrLen := extHdr.ParsedLen() + extHdr.PacketNumber = opener.DecodePacketNumber(extHdr.PacketNumber, extHdr.PacketNumberLen) + decrypted, err := opener.Open(data[extHdrLen:extHdrLen], data[extHdrLen:], extHdr.PacketNumber, data[:extHdrLen]) + if err != nil { + return nil, nil, err + } + if parseErr != nil { + return nil, nil, parseErr + } + return extHdr, decrypted, nil +} + +func (u *packetUnpacker) unpackShortHeaderPacket(opener handshake.ShortHeaderOpener, rcvTime monotime.Time, data []byte) (protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, []byte, error) { + l, pn, pnLen, kp, parseErr := u.unpackShortHeader(opener, data) + // If the reserved bits are set incorrectly, we still need to continue unpacking. + // This avoids a timing side-channel, which otherwise might allow an attacker + // to gain information about the header encryption. + if parseErr != nil && parseErr != wire.ErrInvalidReservedBits { + return 0, 0, 0, nil, &headerParseError{parseErr} + } + pn = opener.DecodePacketNumber(pn, pnLen) + decrypted, err := opener.Open(data[l:l], data[l:], rcvTime, pn, kp, data[:l]) + if err != nil { + return 0, 0, 0, nil, err + } + return pn, pnLen, kp, decrypted, parseErr +} + +func (u *packetUnpacker) unpackShortHeader(hd headerDecryptor, data []byte) (int, protocol.PacketNumber, protocol.PacketNumberLen, protocol.KeyPhaseBit, error) { + hdrLen := 1 /* first header byte */ + u.shortHdrConnIDLen + if len(data) < hdrLen+4+16 { + return 0, 0, 0, 0, fmt.Errorf("packet too small, expected at least 20 bytes after the header, got %d", len(data)-hdrLen) + } + origPNBytes := make([]byte, 4) + copy(origPNBytes, data[hdrLen:hdrLen+4]) + // 2. decrypt the header, assuming a 4 byte packet number + hd.DecryptHeader( + data[hdrLen+4:hdrLen+4+16], + &data[0], + data[hdrLen:hdrLen+4], + ) + // 3. parse the header (and learn the actual length of the packet number) + l, pn, pnLen, kp, parseErr := wire.ParseShortHeader(data, u.shortHdrConnIDLen) + if parseErr != nil && parseErr != wire.ErrInvalidReservedBits { + return l, pn, pnLen, kp, parseErr + } + // 4. if the packet number is shorter than 4 bytes, replace the remaining bytes with the copy we saved earlier + if pnLen != protocol.PacketNumberLen4 { + copy(data[hdrLen+int(pnLen):hdrLen+4], origPNBytes[int(pnLen):]) + } + return l, pn, pnLen, kp, parseErr +} + +// The error is either nil, a wire.ErrInvalidReservedBits or of type headerParseError. +func (u *packetUnpacker) unpackLongHeader(hd headerDecryptor, hdr *wire.Header, data []byte) (*wire.ExtendedHeader, error) { + extHdr, err := unpackLongHeader(hd, hdr, data) + if err != nil && err != wire.ErrInvalidReservedBits { + return nil, &headerParseError{err: err} + } + return extHdr, err +} + +func unpackLongHeader(hd headerDecryptor, hdr *wire.Header, data []byte) (*wire.ExtendedHeader, error) { + hdrLen := hdr.ParsedLen() + if protocol.ByteCount(len(data)) < hdrLen+4+16 { + return nil, fmt.Errorf("packet too small, expected at least 20 bytes after the header, got %d", protocol.ByteCount(len(data))-hdrLen) + } + // The packet number can be up to 4 bytes long, but we won't know the length until we decrypt it. + // 1. save a copy of the 4 bytes + origPNBytes := make([]byte, 4) + copy(origPNBytes, data[hdrLen:hdrLen+4]) + // 2. decrypt the header, assuming a 4 byte packet number + hd.DecryptHeader( + data[hdrLen+4:hdrLen+4+16], + &data[0], + data[hdrLen:hdrLen+4], + ) + // 3. parse the header (and learn the actual length of the packet number) + extHdr, parseErr := hdr.ParseExtended(data) + if parseErr != nil && parseErr != wire.ErrInvalidReservedBits { + return nil, parseErr + } + // 4. if the packet number is shorter than 4 bytes, replace the remaining bytes with the copy we saved earlier + if extHdr.PacketNumberLen != protocol.PacketNumberLen4 { + copy(data[extHdr.ParsedLen():hdrLen+4], origPNBytes[int(extHdr.PacketNumberLen):]) + } + return extHdr, parseErr +} diff --git a/third_party/quic-go/packet_unpacker_test.go b/third_party/quic-go/packet_unpacker_test.go new file mode 100644 index 0000000..d2f69b6 --- /dev/null +++ b/third_party/quic-go/packet_unpacker_test.go @@ -0,0 +1,373 @@ +package quic + +import ( + "crypto/rand" + "testing" + + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/mocks" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type decryptResult struct { + decrypted []byte + err error +} + +func TestUnpackLongHeaderPacket(t *testing.T) { + b := []byte("decrypted") + + t.Run("Initial", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.EncryptionInitial, false, decryptResult{decrypted: b}, nil) + }) + t.Run("Handshake", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.EncryptionHandshake, false, decryptResult{decrypted: b}, nil) + }) + t.Run("0-RTT", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.Encryption0RTT, false, decryptResult{decrypted: b}, nil) + }) +} + +func TestUnpackLongHeaderIncorrectReservedBits(t *testing.T) { + t.Run("decryption fails", func(t *testing.T) { + testUnpackLongHeaderIncorrectReservedBits(t, true) + }) + t.Run("decryption succeeds", func(t *testing.T) { + testUnpackLongHeaderIncorrectReservedBits(t, false) + }) +} + +// Even if the reserved bits are wrong, we still need to continue processing the header. +// This helps prevent a timing side-channel attack, see section 9.5 of RFC 9001. +// We should only return a ErrInvalidReservedBits error if the decryption succeeds, +// as this shows that the peer actually sent an invalid packet. +// However, if decryption fails, this packet is likely injected by an attacker, +// and we should treat it as any other undecryptable packet. +func testUnpackLongHeaderIncorrectReservedBits(t *testing.T, decryptionSucceeds bool) { + decrypted := []byte("decrypted") + expectedErr := wire.ErrInvalidReservedBits + decryptResult := decryptResult{decrypted: decrypted} + if !decryptionSucceeds { + decryptResult.err = handshake.ErrDecryptionFailed + expectedErr = handshake.ErrDecryptionFailed + } + + t.Run("Initial", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.EncryptionInitial, true, decryptResult, expectedErr) + }) + t.Run("Handshake", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.EncryptionHandshake, true, decryptResult, expectedErr) + }) + t.Run("0-RTT", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.Encryption0RTT, true, decryptResult, expectedErr) + }) +} + +func TestUnpackLongHeaderEmptyPayload(t *testing.T) { + expectedErr := &qerr.TransportError{ErrorCode: qerr.ProtocolViolation} + + t.Run("Initial", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.EncryptionInitial, false, decryptResult{}, expectedErr) + }) + + t.Run("Handshake", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.EncryptionHandshake, false, decryptResult{}, expectedErr) + }) + + t.Run("0-RTT", func(t *testing.T) { + testUnpackLongHeaderPacket(t, protocol.Encryption0RTT, false, decryptResult{}, expectedErr) + }) +} + +func testUnpackLongHeaderPacket(t *testing.T, + encLevel protocol.EncryptionLevel, + incorrectReservedBits bool, + decryptResult decryptResult, + expectedErr error, +) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := newPacketUnpacker(cs, 4) + + var packetType protocol.PacketType + switch encLevel { + case protocol.EncryptionInitial: + packetType = protocol.PacketTypeInitial + case protocol.EncryptionHandshake: + packetType = protocol.PacketTypeHandshake + case protocol.Encryption0RTT: + packetType = protocol.PacketType0RTT + } + + payload := []byte("Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.") + extHdr := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: packetType, + Length: protocol.ByteCount(3 + len(payload)), // packet number len + payload + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + Version: protocol.Version1, + }, + PacketNumber: 2, + PacketNumberLen: 3, + } + hdrRaw, err := extHdr.Append(nil, protocol.Version1) + require.NoError(t, err) + if incorrectReservedBits { + hdrRaw[0] |= 0xc + } + data := append(hdrRaw, payload...) + hdr, _, _, err := wire.ParsePacket(data) + require.NoError(t, err) + + opener := mocks.NewMockLongHeaderOpener(mockCtrl) + var calls []any + switch encLevel { + case protocol.EncryptionInitial: + calls = append(calls, cs.EXPECT().GetInitialOpener().Return(opener, nil)) + case protocol.EncryptionHandshake: + calls = append(calls, cs.EXPECT().GetHandshakeOpener().Return(opener, nil)) + case protocol.Encryption0RTT: + calls = append(calls, cs.EXPECT().Get0RTTOpener().Return(opener, nil)) + } + calls = append(calls, []any{ + opener.EXPECT().DecryptHeader(gomock.Any(), gomock.Any(), gomock.Any()), + opener.EXPECT().DecodePacketNumber(protocol.PacketNumber(2), protocol.PacketNumberLen3).Return(protocol.PacketNumber(1234)), + opener.EXPECT().Open(gomock.Any(), payload, protocol.PacketNumber(1234), hdrRaw).Return( + decryptResult.decrypted, decryptResult.err, + ), + }...) + gomock.InOrder(calls...) + + packet, err := unpacker.UnpackLongHeader(hdr, data) + if expectedErr != nil { + require.ErrorIs(t, err, expectedErr) + return + } + require.NoError(t, err) + require.Equal(t, encLevel, packet.encryptionLevel) + require.Equal(t, decryptResult.decrypted, packet.data) +} + +func TestUnpackShortHeaderPacket(t *testing.T) { + testUnpackShortHeaderPacket(t, false, decryptResult{decrypted: []byte("decrypted")}, nil) +} + +func TestUnpackShortHeaderEmptyPayload(t *testing.T) { + testUnpackShortHeaderPacket(t, false, decryptResult{}, &qerr.TransportError{ErrorCode: qerr.ProtocolViolation}) +} + +// Even if the reserved bits are wrong, we still need to continue processing the header. +// This helps prevent a timing side-channel attack, see section 9.5 of RFC 9001. +// We should only return a ErrInvalidReservedBits error if the decryption succeeds, +// as this shows that the peer actually sent an invalid packet. +// However, if decryption fails, this packet is likely injected by an attacker, +// and we should treat it as any other undecryptable packet. +func TestUnpackShortHeaderIncorrectReservedBits(t *testing.T) { + t.Run("decryption fails", func(t *testing.T) { + testUnpackShortHeaderPacket(t, + true, + decryptResult{err: handshake.ErrDecryptionFailed}, + handshake.ErrDecryptionFailed, + ) + }) + + t.Run("decryption succeeds", func(t *testing.T) { + testUnpackShortHeaderPacket(t, + true, + decryptResult{decrypted: []byte("decrypted")}, + wire.ErrInvalidReservedBits, + ) + }) +} + +func testUnpackShortHeaderPacket(t *testing.T, incorrectReservedBits bool, decryptResult decryptResult, expectedErr error) { + mockCtrl := gomock.NewController(t) + connID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5}) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := newPacketUnpacker(cs, connID.Len()) + payload := []byte("Lorem ipsum dolor sit amet") + + hdrRaw, err := wire.AppendShortHeader( + nil, + connID, + 0x1337, + protocol.PacketNumberLen3, + protocol.KeyPhaseOne, + ) + require.NoError(t, err) + if incorrectReservedBits { + hdrRaw[0] |= 0x18 + } + opener := mocks.NewMockShortHeaderOpener(mockCtrl) + opener.EXPECT().DecryptHeader(gomock.Any(), gomock.Any(), gomock.Any()) + cs.EXPECT().Get1RTTOpener().Return(opener, nil) + opener.EXPECT().DecodePacketNumber(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(1234)) + opener.EXPECT().Open(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return( + decryptResult.decrypted, decryptResult.err, + ) + pn, pnLen, kp, data, err := unpacker.UnpackShortHeader(monotime.Now(), append(hdrRaw, payload...)) + if expectedErr != nil { + require.ErrorIs(t, err, expectedErr) + return + } + require.NoError(t, err) + require.Equal(t, decryptResult.decrypted, data) + require.Equal(t, protocol.PacketNumber(1234), pn) + require.Equal(t, protocol.PacketNumberLen3, pnLen) + require.Equal(t, protocol.KeyPhaseOne, kp) +} + +func TestUnpackHeaderSampleLongHeader(t *testing.T) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := newPacketUnpacker(cs, 4) + + extHdr := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + Version: protocol.Version1, + }, + PacketNumber: 1337, + PacketNumberLen: protocol.PacketNumberLen2, + } + data, err := extHdr.Append(nil, protocol.Version1) + require.NoError(t, err) + b := make([]byte, 2+16) // 2 bytes to fill up the packet number, 16 bytes for the sample + rand.Read(b) + data = append(data, b...) + hdr, _, _, err := wire.ParsePacket(data) + require.NoError(t, err) + + t.Run("too short", func(t *testing.T) { + cs.EXPECT().GetHandshakeOpener().Return(mocks.NewMockLongHeaderOpener(mockCtrl), nil) + _, err = unpacker.UnpackLongHeader(hdr, data[:len(data)-1]) + require.IsType(t, &headerParseError{}, err) + require.ErrorContains(t, err, "packet too small, expected at least 20 bytes after the header, got 19") + }) + + t.Run("minimal size", func(t *testing.T) { + opener := mocks.NewMockLongHeaderOpener(mockCtrl) + cs.EXPECT().GetHandshakeOpener().Return(opener, nil) + opener.EXPECT().DecryptHeader(b[len(b)-16:], gomock.Any(), gomock.Any()) + opener.EXPECT().DecodePacketNumber(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(1337)) + opener.EXPECT().Open(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]byte("decrypted"), nil) + _, err = unpacker.UnpackLongHeader(hdr, data) + require.NoError(t, err) + }) +} + +func TestUnpackHeaderSampleShortHeader(t *testing.T) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := newPacketUnpacker(cs, 4) + + data, err := wire.AppendShortHeader( + nil, + protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + 1337, + protocol.PacketNumberLen2, + protocol.KeyPhaseOne, + ) + require.NoError(t, err) + b := make([]byte, 2+16) // 2 bytes to fill up the packet number, 16 bytes for the sample + rand.Read(b) + data = append(data, b...) + + t.Run("too short", func(t *testing.T) { + cs.EXPECT().Get1RTTOpener().Return(mocks.NewMockShortHeaderOpener(mockCtrl), nil) + _, _, _, _, err = unpacker.UnpackShortHeader(monotime.Now(), data[:len(data)-1]) + require.IsType(t, &headerParseError{}, err) + require.ErrorContains(t, err, "packet too small, expected at least 20 bytes after the header, got 19") + }) + + t.Run("minimal size", func(t *testing.T) { + opener := mocks.NewMockShortHeaderOpener(mockCtrl) + cs.EXPECT().Get1RTTOpener().Return(opener, nil) + opener.EXPECT().DecryptHeader(data[len(data)-16:], gomock.Any(), gomock.Any()) + opener.EXPECT().DecodePacketNumber(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(1337)) + opener.EXPECT().Open(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]byte("decrypted"), nil) + _, _, _, _, err = unpacker.UnpackShortHeader(monotime.Now(), data) + require.NoError(t, err) + }) +} + +func TestUnpackErrors(t *testing.T) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := newPacketUnpacker(cs, 4) + + // opener not available + cs.EXPECT().GetHandshakeOpener().Return(nil, handshake.ErrKeysNotYetAvailable) + _, err := unpacker.UnpackLongHeader(&wire.Header{Type: protocol.PacketTypeHandshake}, []byte("foobar")) + require.ErrorIs(t, err, handshake.ErrKeysNotYetAvailable) + + // opener returns error + opener := mocks.NewMockLongHeaderOpener(mockCtrl) + cs.EXPECT().GetHandshakeOpener().Return(opener, nil) + opener.EXPECT().DecryptHeader(gomock.Any(), gomock.Any(), gomock.Any()) + opener.EXPECT().DecodePacketNumber(gomock.Any(), gomock.Any()).Return(protocol.PacketNumber(1234)) + opener.EXPECT().Open(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, &qerr.TransportError{ErrorCode: qerr.CryptoBufferExceeded}) + _, err = unpacker.UnpackLongHeader(&wire.Header{Type: protocol.PacketTypeHandshake}, make([]byte, 100)) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.CryptoBufferExceeded}) +} + +func TestUnpackHeaderDecryption(t *testing.T) { + mockCtrl := gomock.NewController(t) + cs := mocks.NewMockCryptoSetup(mockCtrl) + unpacker := newPacketUnpacker(cs, 4) + connID := protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}) + + extHdr := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + Length: 2, // packet number len + DestConnectionID: connID, + Version: protocol.Version1, + }, + PacketNumber: 0x1337, + PacketNumberLen: protocol.PacketNumberLen2, + } + hdrRaw, err := extHdr.Append(nil, protocol.Version1) + require.NoError(t, err) + hdr, _, _, err := wire.ParsePacket(hdrRaw) + require.NoError(t, err) + + origHdrRaw := append([]byte{}, hdrRaw...) // save a copy of the header + firstHdrByte := hdrRaw[0] + hdrRaw[0] ^= 0xff // invert the first byte + hdrRaw[len(hdrRaw)-2] ^= 0xff // invert the packet number + hdrRaw[len(hdrRaw)-1] ^= 0xff // invert the packet number + require.NotEqual(t, hdrRaw[0], firstHdrByte) + + opener := mocks.NewMockLongHeaderOpener(mockCtrl) + cs.EXPECT().GetHandshakeOpener().Return(opener, nil) + gomock.InOrder( + // we're using a 2 byte packet number, so the sample starts at the 3rd payload byte + opener.EXPECT().DecryptHeader( + []byte{3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18}, + &hdrRaw[0], + append(hdrRaw[len(hdrRaw)-2:], []byte{1, 2}...)).Do(func(_ []byte, firstByte *byte, pnBytes []byte) { + *firstByte ^= 0xff // invert the first byte back + for i := range pnBytes { + pnBytes[i] ^= 0xff // invert the packet number bytes + } + }), + opener.EXPECT().DecodePacketNumber(protocol.PacketNumber(0x1337), protocol.PacketNumberLen2).Return(protocol.PacketNumber(0x7331)), + opener.EXPECT().Open(gomock.Any(), gomock.Any(), protocol.PacketNumber(0x7331), origHdrRaw).Return([]byte{0}, nil), + ) + + data := hdrRaw + for i := 1; i <= 100; i++ { + data = append(data, uint8(i)) + } + packet, err := unpacker.UnpackLongHeader(hdr, data) + require.NoError(t, err) + require.Equal(t, protocol.PacketNumber(0x7331), packet.hdr.PacketNumber) +} diff --git a/third_party/quic-go/path_manager.go b/third_party/quic-go/path_manager.go new file mode 100644 index 0000000..0ba208e --- /dev/null +++ b/third_party/quic-go/path_manager.go @@ -0,0 +1,206 @@ +package quic + +import ( + "crypto/rand" + "net" + "slices" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" +) + +type pathID int64 + +const invalidPathID pathID = -1 + +// Maximum number of paths to keep track of. +// If the peer probes another path (before the pathTimeout of an existing path expires), +// this probing attempt is ignored. +const maxPaths = 3 + +// If no packet is received for a path for pathTimeout, +// the path can be evicted when the peer probes another path. +// This prevents an attacker from churning through paths by duplicating packets and +// sending them with spoofed source addresses. +const pathTimeout = 5 * time.Second + +type path struct { + id pathID + addr net.Addr + lastPacketTime monotime.Time + pathChallenge [8]byte + validated bool + rcvdNonProbing bool +} + +type pathManager struct { + nextPathID pathID + // ordered by lastPacketTime, with the most recently used path at the end + paths []*path + + getConnID func(pathID) (_ protocol.ConnectionID, ok bool) + retireConnID func(pathID) + + logger utils.Logger +} + +func newPathManager( + getConnID func(pathID) (_ protocol.ConnectionID, ok bool), + retireConnID func(pathID), + logger utils.Logger, +) *pathManager { + return &pathManager{ + paths: make([]*path, 0, maxPaths+1), + getConnID: getConnID, + retireConnID: retireConnID, + logger: logger, + } +} + +// Returns a path challenge frame if one should be sent. +// May return nil. +func (pm *pathManager) HandlePacket( + remoteAddr net.Addr, + t monotime.Time, + pathChallenge *wire.PathChallengeFrame, // may be nil if the packet didn't contain a PATH_CHALLENGE + isNonProbing bool, +) (_ protocol.ConnectionID, _ []ackhandler.Frame, shouldSwitch bool) { + var p *path + for i, path := range pm.paths { + if addrsEqual(path.addr, remoteAddr) { + p = path + p.lastPacketTime = t + // already sent a PATH_CHALLENGE for this path + if isNonProbing { + path.rcvdNonProbing = true + } + if pm.logger.Debug() { + pm.logger.Debugf("received packet for path %s that was already probed, validated: %t", remoteAddr, path.validated) + } + shouldSwitch = path.validated && path.rcvdNonProbing + if i != len(pm.paths)-1 { + // move the path to the end of the list + pm.paths = slices.Delete(pm.paths, i, i+1) + pm.paths = append(pm.paths, p) + } + if pathChallenge == nil { + return protocol.ConnectionID{}, nil, shouldSwitch + } + } + } + + if len(pm.paths) >= maxPaths { + if pm.paths[0].lastPacketTime.Add(pathTimeout).After(t) { + if pm.logger.Debug() { + pm.logger.Debugf("received packet for previously unseen path %s, but already have %d paths", remoteAddr, len(pm.paths)) + } + return protocol.ConnectionID{}, nil, shouldSwitch + } + // evict the oldest path, if the last packet was received more than pathTimeout ago + pm.retireConnID(pm.paths[0].id) + pm.paths = pm.paths[1:] + } + + var pathID pathID + if p != nil { + pathID = p.id + } else { + pathID = pm.nextPathID + } + + // previously unseen path, initiate path validation by sending a PATH_CHALLENGE + connID, ok := pm.getConnID(pathID) + if !ok { + pm.logger.Debugf("skipping validation of new path %s since no connection ID is available", remoteAddr) + return protocol.ConnectionID{}, nil, shouldSwitch + } + + frames := make([]ackhandler.Frame, 0, 2) + if p == nil { + var pathChallengeData [8]byte + rand.Read(pathChallengeData[:]) + p = &path{ + id: pm.nextPathID, + addr: remoteAddr, + lastPacketTime: t, + rcvdNonProbing: isNonProbing, + pathChallenge: pathChallengeData, + } + pm.nextPathID++ + pm.paths = append(pm.paths, p) + frames = append(frames, ackhandler.Frame{ + Frame: &wire.PathChallengeFrame{Data: p.pathChallenge}, + Handler: (*pathManagerAckHandler)(pm), + }) + pm.logger.Debugf("enqueueing PATH_CHALLENGE for new path %s", remoteAddr) + } + if pathChallenge != nil { + frames = append(frames, ackhandler.Frame{ + Frame: &wire.PathResponseFrame{Data: pathChallenge.Data}, + Handler: (*pathManagerAckHandler)(pm), + }) + } + return connID, frames, shouldSwitch +} + +func (pm *pathManager) HandlePathResponseFrame(f *wire.PathResponseFrame) { + for _, p := range pm.paths { + if f.Data == p.pathChallenge { + // path validated + p.validated = true + pm.logger.Debugf("path %s validated", p.addr) + break + } + } +} + +// SwitchToPath is called when the connection switches to a new path +func (pm *pathManager) SwitchToPath(addr net.Addr) { + // retire all other paths + for _, path := range pm.paths { + if addrsEqual(path.addr, addr) { + pm.logger.Debugf("switching to path %d (%s)", path.id, addr) + continue + } + pm.retireConnID(path.id) + } + clear(pm.paths) + pm.paths = pm.paths[:0] +} + +type pathManagerAckHandler pathManager + +var _ ackhandler.FrameHandler = &pathManagerAckHandler{} + +// Acknowledging the frame doesn't validate the path, only receiving the PATH_RESPONSE does. +func (pm *pathManagerAckHandler) OnAcked(f wire.Frame) {} + +func (pm *pathManagerAckHandler) OnLost(f wire.Frame) { + pc, ok := f.(*wire.PathChallengeFrame) + if !ok { + return + } + for i, path := range pm.paths { + if path.pathChallenge == pc.Data { + pm.paths = slices.Delete(pm.paths, i, i+1) + pm.retireConnID(path.id) + break + } + } +} + +func addrsEqual(addr1, addr2 net.Addr) bool { + if addr1 == nil || addr2 == nil { + return false + } + a1, ok1 := addr1.(*net.UDPAddr) + a2, ok2 := addr2.(*net.UDPAddr) + if ok1 && ok2 { + return a1.IP.Equal(a2.IP) && a1.Port == a2.Port + } + return addr1.String() == addr2.String() +} diff --git a/third_party/quic-go/path_manager_outgoing.go b/third_party/quic-go/path_manager_outgoing.go new file mode 100644 index 0000000..2c6e41e --- /dev/null +++ b/third_party/quic-go/path_manager_outgoing.go @@ -0,0 +1,314 @@ +package quic + +import ( + "context" + "crypto/rand" + "errors" + "slices" + "sync" + "sync/atomic" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +var ( + // ErrPathClosed is returned when trying to switch to a path that has been closed. + ErrPathClosed = errors.New("path closed") + // ErrPathNotValidated is returned when trying to use a path before path probing has completed. + ErrPathNotValidated = errors.New("path not yet validated") +) + +var errPathDoesNotExist = errors.New("path does not exist") + +// Path is a network path. +type Path struct { + id pathID + pathManager *pathManagerOutgoing + tr *Transport + initialRTT time.Duration + + enablePath func() + validated atomic.Bool + abandon chan struct{} +} + +func (p *Path) Probe(ctx context.Context) error { + path := p.pathManager.addPath(p, p.enablePath) + + p.pathManager.enqueueProbe(p) + nextProbeDur := p.initialRTT + var timer *time.Timer + var timerChan <-chan time.Time + for { + select { + case <-ctx.Done(): + return context.Cause(ctx) + case <-path.Validated(): + p.validated.Store(true) + return nil + case <-timerChan: + nextProbeDur *= 2 // exponential backoff + p.pathManager.enqueueProbe(p) + case <-path.ProbeSent(): + case <-p.abandon: + return ErrPathClosed + } + + if timer != nil { + timer.Stop() + } + timer = time.NewTimer(nextProbeDur) + timerChan = timer.C + } +} + +// Switch switches the QUIC connection to this path. +// It immediately stops sending on the old path, and sends on this new path. +func (p *Path) Switch() error { + if err := p.pathManager.switchToPath(p.id); err != nil { + switch { + case errors.Is(err, ErrPathNotValidated): + return err + case errors.Is(err, errPathDoesNotExist) && !p.validated.Load(): + select { + case <-p.abandon: + return ErrPathClosed + default: + return ErrPathNotValidated + } + default: + return ErrPathClosed + } + } + return nil +} + +// Close abandons a path. +// It is not possible to close the path that’s currently active. +// After closing, it is not possible to probe this path again. +func (p *Path) Close() error { + select { + case <-p.abandon: + return nil + default: + } + + if err := p.pathManager.removePath(p.id); err != nil { + return err + } + close(p.abandon) + return nil +} + +type pathOutgoing struct { + pathChallenges [][8]byte // length is implicitly limited by exponential backoff + tr *Transport + isValidated bool + probeSent chan struct{} // receives when a PATH_CHALLENGE is sent + validated chan struct{} // closed when the path the corresponding PATH_RESPONSE is received + enablePath func() +} + +func (p *pathOutgoing) ProbeSent() <-chan struct{} { return p.probeSent } +func (p *pathOutgoing) Validated() <-chan struct{} { return p.validated } + +type pathManagerOutgoing struct { + getConnID func(pathID) (_ protocol.ConnectionID, ok bool) + retireConnID func(pathID) + scheduleSending func() + + mx sync.Mutex + activePath pathID + pathsToProbe []pathID + paths map[pathID]*pathOutgoing + nextPathID pathID + pathToSwitchTo *pathOutgoing +} + +// newPathManagerOutgoing creates a new pathManagerOutgoing object. This +// function must be side-effect free as it may be called multiple times for a +// single connection. +func newPathManagerOutgoing( + getConnID func(pathID) (_ protocol.ConnectionID, ok bool), + retireConnID func(pathID), + scheduleSending func(), +) *pathManagerOutgoing { + return &pathManagerOutgoing{ + activePath: 0, // at initialization time, we're guaranteed to be using the handshake path + nextPathID: 1, + getConnID: getConnID, + retireConnID: retireConnID, + scheduleSending: scheduleSending, + paths: make(map[pathID]*pathOutgoing, 4), + } +} + +func (pm *pathManagerOutgoing) addPath(p *Path, enablePath func()) *pathOutgoing { + pm.mx.Lock() + defer pm.mx.Unlock() + + // path might already exist, and just being re-probed + if existingPath, ok := pm.paths[p.id]; ok { + existingPath.validated = make(chan struct{}) + return existingPath + } + + path := &pathOutgoing{ + tr: p.tr, + probeSent: make(chan struct{}, 1), + validated: make(chan struct{}), + enablePath: enablePath, + } + pm.paths[p.id] = path + return path +} + +func (pm *pathManagerOutgoing) enqueueProbe(p *Path) { + pm.mx.Lock() + pm.pathsToProbe = append(pm.pathsToProbe, p.id) + pm.mx.Unlock() + pm.scheduleSending() +} + +func (pm *pathManagerOutgoing) removePath(id pathID) error { + if err := pm.removePathImpl(id); err != nil { + return err + } + pm.scheduleSending() + return nil +} + +func (pm *pathManagerOutgoing) removePathImpl(id pathID) error { + pm.mx.Lock() + defer pm.mx.Unlock() + + if id == pm.activePath { + return errors.New("cannot close active path") + } + p, ok := pm.paths[id] + if !ok { + return nil + } + if len(p.pathChallenges) > 0 { + pm.retireConnID(id) + } + delete(pm.paths, id) + return nil +} + +func (pm *pathManagerOutgoing) switchToPath(id pathID) error { + pm.mx.Lock() + defer pm.mx.Unlock() + + p, ok := pm.paths[id] + if !ok { + return errPathDoesNotExist + } + if !p.isValidated { + return ErrPathNotValidated + } + pm.pathToSwitchTo = p + pm.activePath = id + return nil +} + +func (pm *pathManagerOutgoing) NewPath(t *Transport, initialRTT time.Duration, enablePath func()) *Path { + pm.mx.Lock() + defer pm.mx.Unlock() + + id := pm.nextPathID + pm.nextPathID++ + return &Path{ + pathManager: pm, + id: id, + tr: t, + enablePath: enablePath, + initialRTT: initialRTT, + abandon: make(chan struct{}), + } +} + +func (pm *pathManagerOutgoing) NextPathToProbe() (_ protocol.ConnectionID, _ ackhandler.Frame, _ *Transport, hasPath bool) { + pm.mx.Lock() + defer pm.mx.Unlock() + + var p *pathOutgoing + id := invalidPathID + for _, pID := range pm.pathsToProbe { + var ok bool + p, ok = pm.paths[pID] + if ok { + id = pID + break + } + // if the path doesn't exist in the map, it might have been abandoned + pm.pathsToProbe = pm.pathsToProbe[1:] + } + if id == invalidPathID { + return protocol.ConnectionID{}, ackhandler.Frame{}, nil, false + } + + connID, ok := pm.getConnID(id) + if !ok { + return protocol.ConnectionID{}, ackhandler.Frame{}, nil, false + } + + var b [8]byte + _, _ = rand.Read(b[:]) + p.pathChallenges = append(p.pathChallenges, b) + + pm.pathsToProbe = pm.pathsToProbe[1:] + p.enablePath() + select { + case p.probeSent <- struct{}{}: + default: + } + frame := ackhandler.Frame{ + Frame: &wire.PathChallengeFrame{Data: b}, + Handler: (*pathManagerOutgoingAckHandler)(pm), + } + return connID, frame, p.tr, true +} + +func (pm *pathManagerOutgoing) HandlePathResponseFrame(f *wire.PathResponseFrame) { + pm.mx.Lock() + defer pm.mx.Unlock() + + for _, p := range pm.paths { + if slices.Contains(p.pathChallenges, f.Data) { + // path validated + if !p.isValidated { + // make sure that duplicate PATH_RESPONSE frames are ignored + p.isValidated = true + p.pathChallenges = nil + close(p.validated) + } + break + } + } +} + +func (pm *pathManagerOutgoing) ShouldSwitchPath() (*Transport, bool) { + pm.mx.Lock() + defer pm.mx.Unlock() + + if pm.pathToSwitchTo == nil { + return nil, false + } + p := pm.pathToSwitchTo + pm.pathToSwitchTo = nil + return p.tr, true +} + +type pathManagerOutgoingAckHandler pathManagerOutgoing + +var _ ackhandler.FrameHandler = &pathManagerOutgoingAckHandler{} + +// OnAcked is called when the PATH_CHALLENGE is acked. +// This doesn't validate the path, only receiving the PATH_RESPONSE does. +func (pm *pathManagerOutgoingAckHandler) OnAcked(wire.Frame) {} + +func (pm *pathManagerOutgoingAckHandler) OnLost(wire.Frame) {} diff --git a/third_party/quic-go/path_manager_outgoing_test.go b/third_party/quic-go/path_manager_outgoing_test.go new file mode 100644 index 0000000..8938ad9 --- /dev/null +++ b/third_party/quic-go/path_manager_outgoing_test.go @@ -0,0 +1,287 @@ +package quic + +import ( + "context" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestPathManagerOutgoingPathProbing(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + connIDs := []protocol.ConnectionID{ + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + } + pm := newPathManagerOutgoing( + func(id pathID) (protocol.ConnectionID, bool) { + connID := connIDs[0] + connIDs = connIDs[1:] + return connID, true + }, + func(id pathID) { t.Fatal("didn't expect any connection ID to be retired") }, + func() {}, + ) + + _, _, _, ok := pm.NextPathToProbe() + require.False(t, ok) + + tr1 := &Transport{} + var enabled bool + p := pm.NewPath(tr1, time.Second, func() { enabled = true }) + require.ErrorIs(t, p.Switch(), ErrPathNotValidated) + + errChan := make(chan error, 1) + go func() { errChan <- p.Probe(context.Background()) }() + + // wait for the path to be queued for probing + synctest.Wait() + + require.False(t, enabled) + connID, f, tr, ok := pm.NextPathToProbe() + require.True(t, ok) + require.Equal(t, tr1, tr) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), connID) + require.IsType(t, &wire.PathChallengeFrame{}, f.Frame) + pc := f.Frame.(*wire.PathChallengeFrame) + require.True(t, enabled) + + _, _, _, ok = pm.NextPathToProbe() + require.False(t, ok) + + select { + case <-errChan: + t.Fatal("should still be probing") + default: + } + + // acking the frame doesn't complete path validation... + f.Handler.OnAcked(f.Frame) + select { + case <-errChan: + t.Fatal("should still be probing") + default: + } + + require.ErrorIs(t, p.Switch(), ErrPathNotValidated) + _, ok = pm.ShouldSwitchPath() + require.False(t, ok) + + // ... neither does receiving a random PATH_RESPONSE... + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: [8]byte{'f', 'o', 'o', 'f', 'o', 'o'}}) + f.Handler.OnAcked(f.Frame) // doesn't do anything + f.Handler.OnLost(f.Frame) // doesn't do anything + select { + case <-errChan: + t.Fatal("should still be probing") + default: + } + + // ... only receiving the corresponding PATH_RESPONSE does + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: pc.Data}) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + + // receiving it multiple times is ok + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: pc.Data}) + + // now switch to the other path + _, ok = pm.ShouldSwitchPath() + require.False(t, ok) + require.NoError(t, p.Switch()) + // the active path can't be closed + require.EqualError(t, p.Close(), "cannot close active path") + switchToTransport, ok := pm.ShouldSwitchPath() + require.True(t, ok) + require.Equal(t, tr1, switchToTransport) + }) +} + +func TestPathManagerOutgoingRetransmissions(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + connIDs := []protocol.ConnectionID{ + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + protocol.ParseConnectionID([]byte{2, 3, 4, 5, 6, 7, 8, 9}), + } + var retiredConnIDs []protocol.ConnectionID + scheduledSending := make(chan struct{}, 20) + pm := newPathManagerOutgoing( + func(id pathID) (protocol.ConnectionID, bool) { return connIDs[id], true }, + func(id pathID) { retiredConnIDs = append(retiredConnIDs, connIDs[id]) }, + func() { scheduledSending <- struct{}{} }, + ) + + _, _, _, ok := pm.NextPathToProbe() + require.False(t, ok) + + tr1 := &Transport{} + const initialRTT = 5 * time.Millisecond + p := pm.NewPath(tr1, initialRTT, func() {}) + + pathChallengeChan := make(chan [8]byte) + done := make(chan struct{}) + defer close(done) + go func() { + for { + select { + case <-scheduledSending: + case <-done: + return + } + _, f, _, ok := pm.NextPathToProbe() + if !ok { + // should never happen + pathChallengeChan <- [8]byte{} + continue + } + pathChallengeChan <- f.Frame.(*wire.PathChallengeFrame).Data + } + }() + + errChan := make(chan error, 1) + go func() { errChan <- p.Probe(context.Background()) }() + + start := time.Now() + type result struct { + pc *[8]byte + took time.Duration + } + var results []result + for range 4 { + select { + case <-errChan: + t.Fatal("probing should not have completed") + case pc := <-pathChallengeChan: + results = append(results, result{pc: &pc, took: time.Since(start)}) + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + + for i, r1 := range results { + require.NotNil(t, r1.pc) + if i > 0 { + took := r1.took - results[i-1].took + t.Log("took", took) + require.Equal(t, took, initialRTT<<(i-1)) + } + for j, r2 := range results { + if i == j { + continue + } + require.NotEqual(t, r1.pc, r2.pc) + } + } + + // receiving a PATH_RESPONSE for any of the PATH_CHALLENGES completes path validation + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: *results[2].pc}) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("probing should have completed") + } + + // It is valid to probe again + results = results[:0] + ctx, cancel := context.WithCancel(context.Background()) + go func() { errChan <- p.Probe(ctx) }() + + synctest.Wait() + + for range 2 { + select { + case err := <-errChan: + require.NoError(t, err) + case pc := <-pathChallengeChan: + results = append(results, result{pc: &pc, took: time.Since(start)}) + case <-time.After(time.Second): + t.Fatal("should have received a path challenge") + } + } + // this time, don't receive a PATH_RESPONSE + cancel() + synctest.Wait() + select { + case err := <-errChan: + require.ErrorIs(t, err, context.Canceled) + default: + t.Fatal("should have received a context canceled error") + } + }) +} + +func TestPathManagerOutgoingAbandonPath(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + connIDs := []protocol.ConnectionID{ + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + } + var retiredPaths []pathID + pm := newPathManagerOutgoing( + func(id pathID) (protocol.ConnectionID, bool) { + connID := connIDs[0] + connIDs = connIDs[1:] + return connID, true + }, + func(id pathID) { retiredPaths = append(retiredPaths, id) }, + func() {}, + ) + + // path abandoned before the PATH_CHALLENGE is sent out + p1 := pm.NewPath(&Transport{}, time.Second, func() {}) + errChan := make(chan error, 1) + go func() { errChan <- p1.Probe(context.Background()) }() + + // wait for the path to be queued for probing + synctest.Wait() + + require.NoError(t, p1.Close()) + // closing the path multiple times is ok + require.NoError(t, p1.Close()) + require.NoError(t, p1.Close()) + _, _, _, ok := pm.NextPathToProbe() + require.False(t, ok) + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, ErrPathClosed) + default: + t.Fatal("should have received a path closed error") + } + require.Empty(t, retiredPaths) + + p2 := pm.NewPath(&Transport{}, time.Second, func() {}) + go func() { errChan <- p2.Probe(context.Background()) }() + + // wait for the path to be queued for probing + synctest.Wait() + + connID, f, _, ok := pm.NextPathToProbe() + require.True(t, ok) + require.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), connID) + + require.NoError(t, p2.Close()) + require.Equal(t, []pathID{p2.id}, retiredPaths) + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: f.Frame.(*wire.PathChallengeFrame).Data}) + _, _, _, ok = pm.NextPathToProbe() + require.False(t, ok) + // it's not possible to switch to an abandoned path + require.ErrorIs(t, p2.Switch(), ErrPathClosed) + }) +} diff --git a/third_party/quic-go/path_manager_test.go b/third_party/quic-go/path_manager_test.go new file mode 100644 index 0000000..7a49710 --- /dev/null +++ b/third_party/quic-go/path_manager_test.go @@ -0,0 +1,356 @@ +package quic + +import ( + "crypto/rand" + "net" + "testing" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +// The path is established by receiving a non-probing packet. +// The first non-probing packet is received after path validation has completed. +// This is the typical scenario when the client initiates connection migration. +func TestPathManagerIntentionalMigration(t *testing.T) { + connIDs := []protocol.ConnectionID{ + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + protocol.ParseConnectionID([]byte{2, 3, 4, 5, 6, 7, 8, 9}), + protocol.ParseConnectionID([]byte{3, 4, 5, 6, 7, 8, 9, 0}), + } + var retiredConnIDs []protocol.ConnectionID + pm := newPathManager( + func(id pathID) (protocol.ConnectionID, bool) { return connIDs[id], true }, + func(id pathID) { retiredConnIDs = append(retiredConnIDs, connIDs[id]) }, + utils.DefaultLogger, + ) + now := monotime.Now() + connID, frames, shouldSwitch := pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + false, + ) + require.Equal(t, connIDs[0], connID) + require.Len(t, frames, 2) + require.IsType(t, &wire.PathChallengeFrame{}, frames[0].Frame) + pc1 := frames[0].Frame.(*wire.PathChallengeFrame) + require.NotZero(t, pc1.Data) + require.NotEqual(t, [8]byte{1, 2, 3, 4, 5, 6, 7, 8}, pc1.Data) + require.IsType(t, &wire.PathResponseFrame{}, frames[1].Frame) + require.Equal(t, [8]byte{1, 2, 3, 4, 5, 6, 7, 8}, frames[1].Frame.(*wire.PathResponseFrame).Data) + require.False(t, shouldSwitch) + + // receiving another packet for the same path doesn't trigger another PATH_CHALLENGE + connID, frames, shouldSwitch = pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + nil, + false, + ) + require.Zero(t, connID) + require.Len(t, frames, 0) + require.False(t, shouldSwitch) + + // receiving a packet for a different path triggers another PATH_CHALLENGE + addr2 := &net.UDPAddr{IP: net.IPv4(5, 6, 7, 8), Port: 1000} + connID, frames, shouldSwitch = pm.HandlePacket(addr2, now, nil, false) + require.Equal(t, connIDs[1], connID) + require.Len(t, frames, 1) + require.IsType(t, &wire.PathChallengeFrame{}, frames[0].Frame) + pc2 := frames[0].Frame.(*wire.PathChallengeFrame) + require.NotEqual(t, pc1.Data, pc2.Data) + require.False(t, shouldSwitch) + + // acknowledging the PATH_CHALLENGE doesn't confirm the path + for _, f := range frames { + f.Handler.OnAcked(f.Frame) + } + connID, frames, shouldSwitch = pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + nil, + false, + ) + require.Zero(t, connID) + require.Empty(t, frames) + require.False(t, shouldSwitch) + + // receiving a PATH_RESPONSE for the second path confirms the path + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: pc2.Data}) + connID, frames, shouldSwitch = pm.HandlePacket(addr2, now, nil, false) + require.Zero(t, connID) + require.Empty(t, frames) + require.False(t, shouldSwitch) // no non-probing packet received yet + require.Empty(t, retiredConnIDs) + + // confirming the path doesn't remove other paths + connID, frames, shouldSwitch = pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + nil, + false, + ) + require.Zero(t, connID) + require.Empty(t, frames) + require.False(t, shouldSwitch) + + // now receive a non-probing packet for the new path + connID, frames, shouldSwitch = pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(5, 6, 7, 8), Port: 1000}, + now, + nil, + true, + ) + require.Zero(t, connID) + require.Empty(t, frames) + require.True(t, shouldSwitch) + + // now switch to the new path + pm.SwitchToPath(&net.UDPAddr{IP: net.IPv4(5, 6, 7, 8), Port: 1000}) + + // switching to the path removes other paths + connID, frames, shouldSwitch = pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, now, nil, false) + require.Equal(t, connIDs[2], connID) + require.NotEmpty(t, frames) + require.NotEqual(t, frames[0].Frame.(*wire.PathChallengeFrame).Data, pc1.Data) + require.False(t, shouldSwitch) + require.Equal(t, []protocol.ConnectionID{connIDs[0]}, retiredConnIDs) +} + +func TestPathManagerMultipleProbes(t *testing.T) { + connIDs := []protocol.ConnectionID{ + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + } + pm := newPathManager( + func(id pathID) (protocol.ConnectionID, bool) { return connIDs[id], true }, + func(id pathID) {}, + utils.DefaultLogger, + ) + now := monotime.Now() + // first receive a packet without a PATH_CHALLENGE + connID, frames, shouldSwitch := pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + nil, + false, + ) + require.Equal(t, connIDs[0], connID) + require.Len(t, frames, 1) + require.IsType(t, &wire.PathChallengeFrame{}, frames[0].Frame) + require.False(t, shouldSwitch) + + // now receive a packet on the same path with a PATH_CHALLENGE + connID, frames, shouldSwitch = pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, + false, + ) + require.Equal(t, connIDs[0], connID) + require.Len(t, frames, 1) + require.Equal(t, &wire.PathResponseFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}}, frames[0].Frame) + require.False(t, shouldSwitch) + + // now receive another packet on the same path with a PATH_RESPONSE + connID, frames, shouldSwitch = pm.HandlePacket( + &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, + now, + &wire.PathChallengeFrame{Data: [8]byte{8, 7, 6, 5, 4, 3, 2, 1}}, + false, + ) + require.Equal(t, connIDs[0], connID) + require.Len(t, frames, 1) + require.Equal(t, &wire.PathResponseFrame{Data: [8]byte{8, 7, 6, 5, 4, 3, 2, 1}}, frames[0].Frame) + require.False(t, shouldSwitch) + + // lose the response packet + frames[0].Handler.OnLost(frames[0].Frame) +} + +// The first packet received on the new path is already a non-probing packet. +// We still need to validate the new path, but we can then switch over immediately. +// This is the typical scenario when a NAT rebinding happens. +func TestPathManagerNATRebinding(t *testing.T) { + connIDs := []protocol.ConnectionID{ + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + } + var retiredConnIDs []protocol.ConnectionID + pm := newPathManager( + func(id pathID) (protocol.ConnectionID, bool) { return connIDs[id], true }, + func(id pathID) { retiredConnIDs = append(retiredConnIDs, connIDs[id]) }, + utils.DefaultLogger, + ) + + now := monotime.Now() + connID, frames, shouldSwitch := pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, now, nil, true) + require.Equal(t, connIDs[0], connID) + require.Len(t, frames, 1) + require.IsType(t, &wire.PathChallengeFrame{}, frames[0].Frame) + pc1 := frames[0].Frame.(*wire.PathChallengeFrame) + require.NotZero(t, pc1.Data) + require.False(t, shouldSwitch) + + // receiving a PATH_RESPONSE for the second path confirms the path + pm.HandlePathResponseFrame(&wire.PathResponseFrame{Data: pc1.Data}) + // we now switch to the new path, as soon as the next packet on that path is received + connID, frames, shouldSwitch = pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}, now, nil, false) + require.Zero(t, connID) + require.Empty(t, frames) + require.True(t, shouldSwitch) +} + +func TestPathManagerLimits(t *testing.T) { + var connIDs []protocol.ConnectionID + for range 2*maxPaths + 2 { + b := make([]byte, 8) + rand.Read(b) + connIDs = append(connIDs, protocol.ParseConnectionID(b)) + } + var retiredConnIDs []protocol.ConnectionID + pm := newPathManager( + func(id pathID) (protocol.ConnectionID, bool) { return connIDs[id], true }, + func(id pathID) { retiredConnIDs = append(retiredConnIDs, connIDs[id]) }, + utils.DefaultLogger, + ) + + now := monotime.Now() + firstPathTime := now + var firstPathConnID protocol.ConnectionID + require.Greater(t, pathTimeout, maxPaths*time.Second) + for i := range maxPaths { + connID, frames, _ := pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000 + i}, now, nil, true) + require.NotEmpty(t, frames) + require.Equal(t, connIDs[i], connID) + if i == 0 { + firstPathConnID = connID + } + now = now.Add(time.Second) + } + // the maximum number of paths is already being probed + now = firstPathTime.Add(pathTimeout).Add(-time.Nanosecond) + connID, frames, _ := pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 2000}, now, nil, true) + require.Zero(t, connID) + require.Empty(t, frames) + + // receiving another packet after the pathTimeout of the first path evicts the first path + now = firstPathTime.Add(pathTimeout) + connIDIndex := maxPaths + connID, frames, _ = pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000 + maxPaths}, now, nil, true) + require.NotEmpty(t, frames) + require.Equal(t, connIDs[connIDIndex], connID) + require.Equal(t, []protocol.ConnectionID{firstPathConnID}, retiredConnIDs) + connIDIndex++ + + // switching to a new path frees is up all paths + var f1 []ackhandler.Frame + pm.SwitchToPath(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1000}) + for i := range maxPaths { + connID, frames, _ := pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 3000 + i}, now, nil, true) + if i == 0 { + f1 = frames + } + require.NotEmpty(t, frames) + require.Equal(t, connIDs[connIDIndex], connID) + connIDIndex++ + } + // again, the maximum number of paths is already being probed + connID, frames, _ = pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 2000}, now, nil, true) + require.Zero(t, connID) + require.Empty(t, frames) + + // losing the frame removes this path + f1[0].Handler.OnLost(f1[0].Frame) + + // we can open exactly one more path + connID, frames, _ = pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 4000}, now, nil, true) + require.NotEmpty(t, frames) + require.Equal(t, connIDs[connIDIndex], connID) + connID, frames, _ = pm.HandlePacket(&net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 4001}, now, nil, true) + require.Zero(t, connID) + require.Empty(t, frames) +} + +type mockAddr struct { + str string +} + +func (a *mockAddr) Network() string { return "mock" } +func (a *mockAddr) String() string { return a.str } + +func TestAddrsEqual(t *testing.T) { + tests := []struct { + name string + addr1 net.Addr + addr2 net.Addr + expected bool + }{ + { + name: "nil addresses", + addr1: nil, + addr2: nil, + expected: false, + }, + { + name: "one nil address", + addr1: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234}, + addr2: nil, + expected: false, + }, + { + name: "same IPv4 addresses", + addr1: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234}, + addr2: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234}, + expected: true, + }, + { + name: "different IPv4 addresses", + addr1: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234}, + addr2: &net.UDPAddr{IP: net.IPv4(4, 3, 2, 1), Port: 1234}, + expected: false, + }, + { + name: "different ports", + addr1: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234}, + addr2: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 4321}, + expected: false, + }, + { + name: "same IPv6 addresses", + addr1: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1234}, + addr2: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1234}, + expected: true, + }, + { + name: "different IPv6 addresses", + addr1: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1234}, + addr2: &net.UDPAddr{IP: net.ParseIP("2001:db8::2"), Port: 1234}, + expected: false, + }, + { + name: "non-UDP addresses with same string representation", + addr1: &mockAddr{str: "192.0.2.1:1234"}, + addr2: &mockAddr{str: "192.0.2.1:1234"}, + expected: true, + }, + { + name: "non-UDP addresses with different string representation", + addr1: &mockAddr{str: "192.0.2.1:1234"}, + addr2: &mockAddr{str: "192.0.2.2:1234"}, + expected: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := addrsEqual(tt.addr1, tt.addr2) + require.Equal(t, tt.expected, result) + }) + } +} diff --git a/third_party/quic-go/qlog/benchmark_test.go b/third_party/quic-go/qlog/benchmark_test.go new file mode 100644 index 0000000..f0ad019 --- /dev/null +++ b/third_party/quic-go/qlog/benchmark_test.go @@ -0,0 +1,86 @@ +package qlog + +import ( + "io" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlogwriter" +) + +type nopWriteCloserImpl struct{ io.Writer } + +func (nopWriteCloserImpl) Close() error { return nil } + +func nopWriteCloser(w io.Writer) io.WriteCloser { + return &nopWriteCloserImpl{Writer: w} +} + +// BenchmarkConnectionTracing aims to benchmark a somewhat realistic connection that sends and receives packets. +func BenchmarkConnectionTracing(b *testing.B) { + b.ReportAllocs() + + srcConnID := protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}) + trace := qlogwriter.NewConnectionFileSeq( + nopWriteCloser(io.Discard), + false, + protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + []string{EventSchema}, + ) + go trace.Run() + tracer := trace.AddProducer() + b.Cleanup(func() { tracer.Close() }) + + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(1337*time.Millisecond, 0) + rttStats.UpdateRTT(1000*time.Millisecond, 10*time.Millisecond) + rttStats.UpdateRTT(800*time.Millisecond, 100*time.Millisecond) + + var i int + for b.Loop() { + i++ + tracer.RecordEvent(&PacketSent{ + Header: PacketHeader{ + PacketType: PacketType1RTT, + PacketNumber: 1234 + protocol.PacketNumber(i), + KeyPhaseBit: KeyPhaseZero, + DestConnectionID: srcConnID, + }, + Raw: RawInfo{Length: 1337}, + ECN: ECT0, + Frames: []Frame{ + {Frame: &AckFrame{AckRanges: []wire.AckRange{{Largest: 12345 + protocol.PacketNumber(2*i), Smallest: 1234 + protocol.PacketNumber(i)}}}}, + {Frame: &MaxStreamDataFrame{StreamID: 42, MaximumStreamData: 987 + protocol.ByteCount(i)}}, + }, + }) + + tracer.RecordEvent(&MetricsUpdated{ + MinRTT: rttStats.MinRTT(), + SmoothedRTT: rttStats.SmoothedRTT(), + LatestRTT: rttStats.LatestRTT(), + RTTVariance: rttStats.MeanDeviation(), + CongestionWindow: int(12345 + protocol.ByteCount(2*i)), + BytesInFlight: int(12345 + protocol.ByteCount(i)), + PacketsInFlight: i, + }) + + if i%2 == 0 { + tracer.RecordEvent(&PacketReceived{ + Header: PacketHeader{ + PacketType: PacketType1RTT, + PacketNumber: 1337 + protocol.PacketNumber(i), + KeyPhaseBit: KeyPhaseOne, + DestConnectionID: srcConnID, + }, + Raw: RawInfo{Length: 1337}, + ECN: ECT0, + Frames: []Frame{ + {Frame: &StreamFrame{StreamID: 123, Offset: int64(1234 + protocol.ByteCount(100*i)), Length: 100, Fin: true}}, + }, + }) + } + } +} diff --git a/third_party/quic-go/qlog/event.go b/third_party/quic-go/qlog/event.go new file mode 100644 index 0000000..1fce2d8 --- /dev/null +++ b/third_party/quic-go/qlog/event.go @@ -0,0 +1,849 @@ +package qlog + +import ( + "fmt" + "net/netip" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +func milliseconds(dur time.Duration) float64 { return float64(dur.Nanoseconds()) / 1e6 } + +type encoderHelper struct { + enc *jsontext.Encoder + err error +} + +func (h *encoderHelper) WriteToken(t jsontext.Token) { + if h.err != nil { + return + } + h.err = h.enc.WriteToken(t) +} + +type versions []Version + +func (v versions) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginArray) + for _, e := range v { + h.WriteToken(jsontext.String(fmt.Sprintf("%x", uint32(e)))) + } + h.WriteToken(jsontext.EndArray) + return h.err +} + +type RawInfo struct { + Length int // full packet length, including header and AEAD authentication tag + PayloadLength int // length of the packet payload, excluding AEAD tag +} + +func (i RawInfo) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("length")) + h.WriteToken(jsontext.Uint(uint64(i.Length))) + if i.PayloadLength != 0 { + h.WriteToken(jsontext.String("payload_length")) + h.WriteToken(jsontext.Uint(uint64(i.PayloadLength))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type PathEndpointInfo struct { + IPv4 netip.AddrPort + IPv6 netip.AddrPort +} + +func (p PathEndpointInfo) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if p.IPv4.IsValid() { + h.WriteToken(jsontext.String("ip_v4")) + h.WriteToken(jsontext.String(p.IPv4.Addr().String())) + h.WriteToken(jsontext.String("port_v4")) + h.WriteToken(jsontext.Int(int64(p.IPv4.Port()))) + } + if p.IPv6.IsValid() { + h.WriteToken(jsontext.String("ip_v6")) + h.WriteToken(jsontext.String(p.IPv6.Addr().String())) + h.WriteToken(jsontext.String("port_v6")) + h.WriteToken(jsontext.Int(int64(p.IPv6.Port()))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type StartedConnection struct { + Local PathEndpointInfo + Remote PathEndpointInfo +} + +func (e StartedConnection) Name() string { return "transport:connection_started" } + +func (e StartedConnection) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("local")) + if err := e.Local.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("remote")) + if err := e.Remote.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type VersionInformation struct { + ClientVersions, ServerVersions []Version + ChosenVersion Version +} + +func (e VersionInformation) Name() string { return "transport:version_information" } + +func (e VersionInformation) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if len(e.ClientVersions) > 0 { + h.WriteToken(jsontext.String("client_versions")) + if err := versions(e.ClientVersions).encode(enc); err != nil { + return err + } + } + if len(e.ServerVersions) > 0 { + h.WriteToken(jsontext.String("server_versions")) + if err := versions(e.ServerVersions).encode(enc); err != nil { + return err + } + } + h.WriteToken(jsontext.String("chosen_version")) + h.WriteToken(jsontext.String(fmt.Sprintf("%x", uint32(e.ChosenVersion)))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type ConnectionClosed struct { + Initiator Initiator + + ConnectionError *TransportErrorCode + ApplicationError *ApplicationErrorCode + + Reason string + + Trigger ConnectionCloseTrigger +} + +func (e ConnectionClosed) Name() string { return "transport:connection_closed" } + +func (e ConnectionClosed) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("initiator")) + h.WriteToken(jsontext.String(string(e.Initiator))) + if e.ConnectionError != nil { + h.WriteToken(jsontext.String("connection_error")) + if e.ConnectionError.IsCryptoError() { + h.WriteToken(jsontext.String(fmt.Sprintf("crypto_error_%#x", uint16(*e.ConnectionError)))) + } else { + switch *e.ConnectionError { + case qerr.NoError: + h.WriteToken(jsontext.String("no_error")) + case qerr.InternalError: + h.WriteToken(jsontext.String("internal_error")) + case qerr.ConnectionRefused: + h.WriteToken(jsontext.String("connection_refused")) + case qerr.FlowControlError: + h.WriteToken(jsontext.String("flow_control_error")) + case qerr.StreamLimitError: + h.WriteToken(jsontext.String("stream_limit_error")) + case qerr.StreamStateError: + h.WriteToken(jsontext.String("stream_state_error")) + case qerr.FinalSizeError: + h.WriteToken(jsontext.String("final_size_error")) + case qerr.FrameEncodingError: + h.WriteToken(jsontext.String("frame_encoding_error")) + case qerr.TransportParameterError: + h.WriteToken(jsontext.String("transport_parameter_error")) + case qerr.ConnectionIDLimitError: + h.WriteToken(jsontext.String("connection_id_limit_error")) + case qerr.ProtocolViolation: + h.WriteToken(jsontext.String("protocol_violation")) + case qerr.InvalidToken: + h.WriteToken(jsontext.String("invalid_token")) + case qerr.ApplicationErrorErrorCode: + h.WriteToken(jsontext.String("application_error")) + case qerr.CryptoBufferExceeded: + h.WriteToken(jsontext.String("crypto_buffer_exceeded")) + case qerr.KeyUpdateError: + h.WriteToken(jsontext.String("key_update_error")) + case qerr.AEADLimitReached: + h.WriteToken(jsontext.String("aead_limit_reached")) + case qerr.NoViablePathError: + h.WriteToken(jsontext.String("no_viable_path")) + default: + h.WriteToken(jsontext.String("unknown")) + h.WriteToken(jsontext.String("error_code")) + h.WriteToken(jsontext.Uint(uint64(*e.ConnectionError))) + } + } + } + if e.ApplicationError != nil { + h.WriteToken(jsontext.String("application_error")) + h.WriteToken(jsontext.String("unknown")) + h.WriteToken(jsontext.String("error_code")) + h.WriteToken(jsontext.Uint(uint64(*e.ApplicationError))) + } + if e.ConnectionError != nil || e.ApplicationError != nil { + h.WriteToken(jsontext.String("reason")) + h.WriteToken(jsontext.String(e.Reason)) + } + if e.Trigger != "" { + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(string(e.Trigger))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type PacketSent struct { + Header PacketHeader + Raw RawInfo + DatagramPayloadChecksum DatagramPayloadChecksum + Frames []Frame + ECN ECN + IsCoalesced bool + Trigger string + SupportedVersions []Version +} + +func (e PacketSent) Name() string { return "transport:packet_sent" } + +func (e PacketSent) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + if e.DatagramPayloadChecksum != 0 { + h.WriteToken(jsontext.String("datagram_payload_checksum")) + h.WriteToken(jsontext.Uint(uint64(e.DatagramPayloadChecksum))) + } + if len(e.Frames) > 0 { + h.WriteToken(jsontext.String("frames")) + if err := frames(e.Frames).encode(enc); err != nil { + return err + } + } + if e.IsCoalesced { + h.WriteToken(jsontext.String("is_coalesced")) + h.WriteToken(jsontext.True) + } + if e.ECN != ECNUnsupported { + h.WriteToken(jsontext.String("ecn")) + h.WriteToken(jsontext.String(string(e.ECN))) + } + if e.Trigger != "" { + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(e.Trigger)) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type PacketReceived struct { + Header PacketHeader + Raw RawInfo + DatagramPayloadChecksum DatagramPayloadChecksum + Frames []Frame + ECN ECN + IsCoalesced bool + Trigger string +} + +func (e PacketReceived) Name() string { return "transport:packet_received" } + +func (e PacketReceived) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + if e.DatagramPayloadChecksum != 0 { + h.WriteToken(jsontext.String("datagram_payload_checksum")) + h.WriteToken(jsontext.Uint(uint64(e.DatagramPayloadChecksum))) + } + if len(e.Frames) > 0 { + h.WriteToken(jsontext.String("frames")) + if err := frames(e.Frames).encode(enc); err != nil { + return err + } + } + if e.IsCoalesced { + h.WriteToken(jsontext.String("is_coalesced")) + h.WriteToken(jsontext.True) + } + if e.ECN != ECNUnsupported { + h.WriteToken(jsontext.String("ecn")) + h.WriteToken(jsontext.String(string(e.ECN))) + } + if e.Trigger != "" { + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(e.Trigger)) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type VersionNegotiationReceived struct { + Header PacketHeaderVersionNegotiation + SupportedVersions []Version +} + +func (e VersionNegotiationReceived) Name() string { return "transport:packet_received" } + +func (e VersionNegotiationReceived) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("supported_versions")) + if err := versions(e.SupportedVersions).encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type VersionNegotiationSent struct { + Header PacketHeaderVersionNegotiation + SupportedVersions []Version +} + +func (e VersionNegotiationSent) Name() string { return "transport:packet_sent" } + +func (e VersionNegotiationSent) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("supported_versions")) + if err := versions(e.SupportedVersions).encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type PacketBuffered struct { + Header PacketHeader + Raw RawInfo + DatagramPayloadChecksum DatagramPayloadChecksum +} + +func (e PacketBuffered) Name() string { return "transport:packet_buffered" } + +func (e PacketBuffered) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + if e.DatagramPayloadChecksum != 0 { + h.WriteToken(jsontext.String("datagram_payload_checksum")) + h.WriteToken(jsontext.Uint(uint64(e.DatagramPayloadChecksum))) + } + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String("keys_unavailable")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// PacketDropped is the transport:packet_dropped event. +type PacketDropped struct { + Header PacketHeader + Raw RawInfo + DatagramPayloadChecksum DatagramPayloadChecksum + Trigger PacketDropReason +} + +func (e PacketDropped) Name() string { return "transport:packet_dropped" } + +func (e PacketDropped) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("raw")) + if err := e.Raw.encode(enc); err != nil { + return err + } + if e.DatagramPayloadChecksum != 0 { + h.WriteToken(jsontext.String("datagram_payload_checksum")) + h.WriteToken(jsontext.Uint(uint64(e.DatagramPayloadChecksum))) + } + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(string(e.Trigger))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type MTUUpdated struct { + Value int + Done bool +} + +func (e MTUUpdated) Name() string { return "recovery:mtu_updated" } + +func (e MTUUpdated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("mtu")) + h.WriteToken(jsontext.Uint(uint64(e.Value))) + h.WriteToken(jsontext.String("done")) + h.WriteToken(jsontext.Bool(e.Done)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// MetricsUpdated logs RTT and congestion metrics as defined in the +// recovery:metrics_updated event. +// The PTO count is logged via PTOCountUpdated. +type MetricsUpdated struct { + MinRTT time.Duration + SmoothedRTT time.Duration + LatestRTT time.Duration + RTTVariance time.Duration + CongestionWindow int + BytesInFlight int + PacketsInFlight int +} + +func (e MetricsUpdated) Name() string { return "recovery:metrics_updated" } + +func (e MetricsUpdated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if e.MinRTT != 0 { + h.WriteToken(jsontext.String("min_rtt")) + h.WriteToken(jsontext.Float(milliseconds(e.MinRTT))) + } + if e.SmoothedRTT != 0 { + h.WriteToken(jsontext.String("smoothed_rtt")) + h.WriteToken(jsontext.Float(milliseconds(e.SmoothedRTT))) + } + if e.LatestRTT != 0 { + h.WriteToken(jsontext.String("latest_rtt")) + h.WriteToken(jsontext.Float(milliseconds(e.LatestRTT))) + } + if e.RTTVariance != 0 { + h.WriteToken(jsontext.String("rtt_variance")) + h.WriteToken(jsontext.Float(milliseconds(e.RTTVariance))) + } + if e.CongestionWindow != 0 { + h.WriteToken(jsontext.String("congestion_window")) + h.WriteToken(jsontext.Uint(uint64(e.CongestionWindow))) + } + if e.BytesInFlight != 0 { + h.WriteToken(jsontext.String("bytes_in_flight")) + h.WriteToken(jsontext.Uint(uint64(e.BytesInFlight))) + } + if e.PacketsInFlight != 0 { + h.WriteToken(jsontext.String("packets_in_flight")) + h.WriteToken(jsontext.Uint(uint64(e.PacketsInFlight))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +// PTOCountUpdated logs the pto_count value of the +// recovery:metrics_updated event. +type PTOCountUpdated struct { + PTOCount uint32 +} + +func (e PTOCountUpdated) Name() string { return "recovery:metrics_updated" } + +func (e PTOCountUpdated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("pto_count")) + h.WriteToken(jsontext.Uint(uint64(e.PTOCount))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type PacketLost struct { + Header PacketHeader + Trigger PacketLossReason +} + +func (e PacketLost) Name() string { return "recovery:packet_lost" } + +func (e PacketLost) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("header")) + if err := e.Header.encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(string(e.Trigger))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type SpuriousLoss struct { + EncryptionLevel protocol.EncryptionLevel + PacketNumber protocol.PacketNumber + PacketReordering uint64 + TimeReordering time.Duration +} + +func (e SpuriousLoss) Name() string { return "recovery:spurious_loss" } + +func (e SpuriousLoss) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("packet_number_space")) + h.WriteToken(jsontext.String(encLevelToPacketNumberSpace(e.EncryptionLevel))) + h.WriteToken(jsontext.String("packet_number")) + h.WriteToken(jsontext.Uint(uint64(e.PacketNumber))) + h.WriteToken(jsontext.String("reordering_packets")) + h.WriteToken(jsontext.Uint(e.PacketReordering)) + h.WriteToken(jsontext.String("reordering_time")) + h.WriteToken(jsontext.Float(milliseconds(e.TimeReordering))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type KeyUpdated struct { + Trigger KeyUpdateTrigger + KeyType KeyType + KeyPhase KeyPhase // only set for 1-RTT keys + // we don't log the keys here, so we don't need `old` and `new`. +} + +func (e KeyUpdated) Name() string { return "security:key_updated" } + +func (e KeyUpdated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(string(e.Trigger))) + h.WriteToken(jsontext.String("key_type")) + h.WriteToken(jsontext.String(string(e.KeyType))) + if e.KeyType == KeyTypeClient1RTT || e.KeyType == KeyTypeServer1RTT { + h.WriteToken(jsontext.String("key_phase")) + h.WriteToken(jsontext.Uint(uint64(e.KeyPhase))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type KeyDiscarded struct { + KeyType KeyType + KeyPhase KeyPhase // only set for 1-RTT keys +} + +func (e KeyDiscarded) Name() string { return "security:key_discarded" } + +func (e KeyDiscarded) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if e.KeyType != KeyTypeClient1RTT && e.KeyType != KeyTypeServer1RTT { + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String("tls")) + } + h.WriteToken(jsontext.String("key_type")) + h.WriteToken(jsontext.String(string(e.KeyType))) + if e.KeyType == KeyTypeClient1RTT || e.KeyType == KeyTypeServer1RTT { + h.WriteToken(jsontext.String("key_phase")) + h.WriteToken(jsontext.Uint(uint64(e.KeyPhase))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type ParametersSet struct { + Restore bool + Initiator Initiator + SentBy protocol.Perspective + OriginalDestinationConnectionID protocol.ConnectionID + InitialSourceConnectionID protocol.ConnectionID + RetrySourceConnectionID *protocol.ConnectionID + StatelessResetToken *protocol.StatelessResetToken + DisableActiveMigration bool + MaxIdleTimeout time.Duration + MaxUDPPayloadSize protocol.ByteCount + AckDelayExponent uint8 + MaxAckDelay time.Duration + ActiveConnectionIDLimit uint64 + InitialMaxData protocol.ByteCount + InitialMaxStreamDataBidiLocal protocol.ByteCount + InitialMaxStreamDataBidiRemote protocol.ByteCount + InitialMaxStreamDataUni protocol.ByteCount + InitialMaxStreamsBidi int64 + InitialMaxStreamsUni int64 + PreferredAddress *PreferredAddress + MaxDatagramFrameSize protocol.ByteCount + EnableResetStreamAt bool +} + +func (e ParametersSet) Name() string { + if e.Restore { + return "transport:parameters_restored" + } + return "transport:parameters_set" +} + +func (e ParametersSet) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if !e.Restore { + h.WriteToken(jsontext.String("initiator")) + h.WriteToken(jsontext.String(string(e.Initiator))) + if e.SentBy == protocol.PerspectiveServer { + h.WriteToken(jsontext.String("original_destination_connection_id")) + h.WriteToken(jsontext.String(e.OriginalDestinationConnectionID.String())) + if e.StatelessResetToken != nil { + h.WriteToken(jsontext.String("stateless_reset_token")) + h.WriteToken(jsontext.String(fmt.Sprintf("%x", e.StatelessResetToken[:]))) + } + if e.RetrySourceConnectionID != nil { + h.WriteToken(jsontext.String("retry_source_connection_id")) + h.WriteToken(jsontext.String((*e.RetrySourceConnectionID).String())) + } + } + h.WriteToken(jsontext.String("initial_source_connection_id")) + h.WriteToken(jsontext.String(e.InitialSourceConnectionID.String())) + } + h.WriteToken(jsontext.String("disable_active_migration")) + h.WriteToken(jsontext.Bool(e.DisableActiveMigration)) + if e.MaxIdleTimeout != 0 { + h.WriteToken(jsontext.String("max_idle_timeout")) + h.WriteToken(jsontext.Float(milliseconds(e.MaxIdleTimeout))) + } + if e.MaxUDPPayloadSize != 0 { + h.WriteToken(jsontext.String("max_udp_payload_size")) + h.WriteToken(jsontext.Int(int64(e.MaxUDPPayloadSize))) + } + if e.AckDelayExponent != 0 { + h.WriteToken(jsontext.String("ack_delay_exponent")) + h.WriteToken(jsontext.Uint(uint64(e.AckDelayExponent))) + } + if e.MaxAckDelay != 0 { + h.WriteToken(jsontext.String("max_ack_delay")) + h.WriteToken(jsontext.Float(milliseconds(e.MaxAckDelay))) + } + if e.ActiveConnectionIDLimit != 0 { + h.WriteToken(jsontext.String("active_connection_id_limit")) + h.WriteToken(jsontext.Uint(e.ActiveConnectionIDLimit)) + } + if e.InitialMaxData != 0 { + h.WriteToken(jsontext.String("initial_max_data")) + h.WriteToken(jsontext.Int(int64(e.InitialMaxData))) + } + if e.InitialMaxStreamDataBidiLocal != 0 { + h.WriteToken(jsontext.String("initial_max_stream_data_bidi_local")) + h.WriteToken(jsontext.Int(int64(e.InitialMaxStreamDataBidiLocal))) + } + if e.InitialMaxStreamDataBidiRemote != 0 { + h.WriteToken(jsontext.String("initial_max_stream_data_bidi_remote")) + h.WriteToken(jsontext.Int(int64(e.InitialMaxStreamDataBidiRemote))) + } + if e.InitialMaxStreamDataUni != 0 { + h.WriteToken(jsontext.String("initial_max_stream_data_uni")) + h.WriteToken(jsontext.Int(int64(e.InitialMaxStreamDataUni))) + } + if e.InitialMaxStreamsBidi != 0 { + h.WriteToken(jsontext.String("initial_max_streams_bidi")) + h.WriteToken(jsontext.Int(e.InitialMaxStreamsBidi)) + } + if e.InitialMaxStreamsUni != 0 { + h.WriteToken(jsontext.String("initial_max_streams_uni")) + h.WriteToken(jsontext.Int(e.InitialMaxStreamsUni)) + } + if e.PreferredAddress != nil { + h.WriteToken(jsontext.String("preferred_address")) + if err := e.PreferredAddress.encode(enc); err != nil { + return err + } + } + if e.MaxDatagramFrameSize != protocol.InvalidByteCount { + h.WriteToken(jsontext.String("max_datagram_frame_size")) + h.WriteToken(jsontext.Int(int64(e.MaxDatagramFrameSize))) + } + if e.EnableResetStreamAt { + h.WriteToken(jsontext.String("reset_stream_at")) + h.WriteToken(jsontext.True) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type PreferredAddress struct { + IPv4, IPv6 netip.AddrPort + ConnectionID protocol.ConnectionID + StatelessResetToken protocol.StatelessResetToken +} + +func (a PreferredAddress) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + if a.IPv4.IsValid() { + h.WriteToken(jsontext.String("ip_v4")) + h.WriteToken(jsontext.String(a.IPv4.Addr().String())) + h.WriteToken(jsontext.String("port_v4")) + h.WriteToken(jsontext.Uint(uint64(a.IPv4.Port()))) + } + if a.IPv6.IsValid() { + h.WriteToken(jsontext.String("ip_v6")) + h.WriteToken(jsontext.String(a.IPv6.Addr().String())) + h.WriteToken(jsontext.String("port_v6")) + h.WriteToken(jsontext.Uint(uint64(a.IPv6.Port()))) + } + h.WriteToken(jsontext.String("connection_id")) + h.WriteToken(jsontext.String(a.ConnectionID.String())) + h.WriteToken(jsontext.String("stateless_reset_token")) + h.WriteToken(jsontext.String(fmt.Sprintf("%x", a.StatelessResetToken))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type LossTimerUpdated struct { + Type LossTimerUpdateType + TimerType TimerType + EncLevel EncryptionLevel + Time time.Time +} + +func (e LossTimerUpdated) Name() string { return "recovery:loss_timer_updated" } + +func (e LossTimerUpdated) Encode(enc *jsontext.Encoder, t time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("event_type")) + h.WriteToken(jsontext.String(string(e.Type))) + h.WriteToken(jsontext.String("timer_type")) + h.WriteToken(jsontext.String(string(e.TimerType))) + h.WriteToken(jsontext.String("packet_number_space")) + h.WriteToken(jsontext.String(encLevelToPacketNumberSpace(e.EncLevel))) + if e.Type == LossTimerUpdateTypeSet { + h.WriteToken(jsontext.String("delta")) + h.WriteToken(jsontext.Float(milliseconds(e.Time.Sub(t)))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type eventLossTimerCanceled struct{} + +func (e eventLossTimerCanceled) Name() string { return "recovery:loss_timer_updated" } + +func (e eventLossTimerCanceled) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("event_type")) + h.WriteToken(jsontext.String("cancelled")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type CongestionStateUpdated struct { + State CongestionState +} + +func (e CongestionStateUpdated) Name() string { return "recovery:congestion_state_updated" } + +func (e CongestionStateUpdated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("new")) + h.WriteToken(jsontext.String(e.State.String())) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type ECNStateUpdated struct { + State ECNState + Trigger string +} + +func (e ECNStateUpdated) Name() string { return "recovery:ecn_state_updated" } + +func (e ECNStateUpdated) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("new")) + h.WriteToken(jsontext.String(string(e.State))) + if e.Trigger != "" { + h.WriteToken(jsontext.String("trigger")) + h.WriteToken(jsontext.String(e.Trigger)) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +type ALPNInformation struct { + ChosenALPN string +} + +func (e ALPNInformation) Name() string { return "transport:alpn_information" } + +func (e ALPNInformation) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("chosen_alpn")) + h.WriteToken(jsontext.String(e.ChosenALPN)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// DebugEvent is a generic event that can be used to log arbitrary messages. +type DebugEvent struct { + EventName string + Message string +} + +func (e DebugEvent) Name() string { + if e.EventName == "" { + return "transport:debug" + } + return fmt.Sprintf("transport:%s", e.EventName) +} + +func (e DebugEvent) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("message")) + h.WriteToken(jsontext.String(e.Message)) + h.WriteToken(jsontext.EndObject) + return h.err +} diff --git a/third_party/quic-go/qlog/event_test.go b/third_party/quic-go/qlog/event_test.go new file mode 100644 index 0000000..ba33f61 --- /dev/null +++ b/third_party/quic-go/qlog/event_test.go @@ -0,0 +1,899 @@ +package qlog + +import ( + "bytes" + "encoding/json" + "net/netip" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlogwriter" + + "github.com/stretchr/testify/require" +) + +func testEventEncoding(t *testing.T, ev qlogwriter.Event) (string, map[string]any) { + t.Helper() + var buf bytes.Buffer + + synctest.Test(t, func(t *testing.T) { + tr := qlogwriter.NewConnectionFileSeq( + nopWriteCloser(&buf), + true, + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + []string{EventSchema}, + ) + go tr.Run() + producer := tr.AddProducer() + + synctest.Wait() + time.Sleep(42 * time.Second) + + producer.RecordEvent(ev) + producer.Close() + }) + + return decode(t, buf.String()) +} + +func decode(t *testing.T, data string) (string, map[string]any) { + t.Helper() + + var result map[string]any + + lines := bytes.Split([]byte(data), []byte{'\n'}) + require.Len(t, lines, 3) // the first line is the trace header, the second line is the event, the third line is empty + require.Empty(t, lines[2]) + require.Equal(t, qlogwriter.RecordSeparator, lines[1][0], "expected record separator at start of line") + require.NoError(t, json.Unmarshal(lines[1][1:], &result)) + require.Equal(t, 42*time.Second, time.Duration(result["time"].(float64)*1e6)*time.Nanosecond) + + return result["name"].(string), result["data"].(map[string]any) +} + +func TestStartedConnection(t *testing.T) { + var localInfo, remoteInfo PathEndpointInfo + localInfo.IPv4 = netip.AddrPortFrom(netip.AddrFrom4([4]byte{192, 168, 13, 37}), 42) + ip, err := netip.ParseAddr("2001:db8::1") + require.NoError(t, err) + remoteInfo.IPv6 = netip.AddrPortFrom(ip, 24) + + name, ev := testEventEncoding(t, &StartedConnection{ + Local: localInfo, + Remote: remoteInfo, + }) + + require.Equal(t, "transport:connection_started", name) + + local, ok := ev["local"].(map[string]any) + require.True(t, ok) + require.Equal(t, "192.168.13.37", local["ip_v4"]) + require.Equal(t, float64(42), local["port_v4"]) + + remote, ok := ev["remote"].(map[string]any) + require.True(t, ok) + require.Equal(t, "2001:db8::1", remote["ip_v6"]) + require.Equal(t, float64(24), remote["port_v6"]) +} + +func TestVersionInformation(t *testing.T) { + name, ev := testEventEncoding(t, &VersionInformation{ChosenVersion: 0x1337}) + + require.Equal(t, "transport:version_information", name) + require.Len(t, ev, 1) + require.Equal(t, "1337", ev["chosen_version"]) +} + +func TestVersionInformationWithNegotiation(t *testing.T) { + name, ev := testEventEncoding(t, &VersionInformation{ + ChosenVersion: 0x1337, + ClientVersions: []Version{1, 2, 3}, + ServerVersions: []Version{4, 5, 6}, + }) + + require.Equal(t, "transport:version_information", name) + require.Len(t, ev, 3) + require.Equal(t, "1337", ev["chosen_version"]) + require.Equal(t, []any{"1", "2", "3"}, ev["client_versions"]) + require.Equal(t, []any{"4", "5", "6"}, ev["server_versions"]) +} + +func TestIdleTimeouts(t *testing.T) { + name, ev := testEventEncoding(t, &ConnectionClosed{ + Initiator: InitiatorLocal, + Trigger: ConnectionCloseTriggerIdleTimeout, + }) + + require.Equal(t, "transport:connection_closed", name) + require.Len(t, ev, 2) + require.Equal(t, "local", ev["initiator"]) + require.Equal(t, "idle_timeout", ev["trigger"]) +} + +func TestReceivedStatelessResetPacket(t *testing.T) { + name, ev := testEventEncoding(t, &ConnectionClosed{ + Initiator: InitiatorRemote, + Trigger: ConnectionCloseTriggerStatelessReset, + }) + + require.Equal(t, "transport:connection_closed", name) + require.Len(t, ev, 2) + require.Equal(t, "remote", ev["initiator"]) + require.Equal(t, "stateless_reset", ev["trigger"]) +} + +func TestVersionNegotiationFailure(t *testing.T) { + name, ev := testEventEncoding(t, &ConnectionClosed{ + Initiator: InitiatorLocal, + Trigger: ConnectionCloseTriggerVersionMismatch, + }) + + require.Equal(t, "transport:connection_closed", name) + require.Len(t, ev, 2) + require.Equal(t, "local", ev["initiator"]) + require.Equal(t, "version_mismatch", ev["trigger"]) +} + +func TestApplicationErrors(t *testing.T) { + code := qerr.ApplicationErrorCode(1337) + name, ev := testEventEncoding(t, &ConnectionClosed{ + Initiator: InitiatorRemote, + ApplicationError: &code, + Reason: "foobar", + }) + + require.Equal(t, "transport:connection_closed", name) + require.Len(t, ev, 4) + require.Equal(t, "remote", ev["initiator"]) + require.Equal(t, "unknown", ev["application_error"]) + require.Equal(t, float64(1337), ev["error_code"]) + require.Equal(t, "foobar", ev["reason"]) +} + +func TestTransportErrors(t *testing.T) { + tests := []struct { + code qerr.TransportErrorCode + want string + }{ + {qerr.NoError, "no_error"}, + {qerr.InternalError, "internal_error"}, + {qerr.ConnectionRefused, "connection_refused"}, + {qerr.FlowControlError, "flow_control_error"}, + {qerr.StreamLimitError, "stream_limit_error"}, + {qerr.StreamStateError, "stream_state_error"}, + {qerr.FinalSizeError, "final_size_error"}, + {qerr.FrameEncodingError, "frame_encoding_error"}, + {qerr.TransportParameterError, "transport_parameter_error"}, + {qerr.ConnectionIDLimitError, "connection_id_limit_error"}, + {qerr.ProtocolViolation, "protocol_violation"}, + {qerr.InvalidToken, "invalid_token"}, + {qerr.ApplicationErrorErrorCode, "application_error"}, + {qerr.CryptoBufferExceeded, "crypto_buffer_exceeded"}, + {qerr.KeyUpdateError, "key_update_error"}, + {qerr.AEADLimitReached, "aead_limit_reached"}, + {qerr.NoViablePathError, "no_viable_path"}, + } + + for _, tt := range tests { + t.Run(tt.want, func(t *testing.T) { + code := tt.code + name, ev := testEventEncoding(t, &ConnectionClosed{ + Initiator: InitiatorLocal, + ConnectionError: &code, + Reason: "foobar", + }) + + require.Equal(t, "transport:connection_closed", name) + require.Equal(t, "local", ev["initiator"]) + require.Equal(t, tt.want, ev["connection_error"]) + require.Equal(t, "foobar", ev["reason"]) + require.NotContains(t, ev, "error_code") + }) + } +} + +func TestTransportCryptoError(t *testing.T) { + code := qerr.TransportErrorCode(0x100 + 0x2a) + name, ev := testEventEncoding(t, &ConnectionClosed{ + Initiator: InitiatorLocal, + ConnectionError: &code, + Reason: "foobar", + }) + + require.Equal(t, "transport:connection_closed", name) + require.Equal(t, "local", ev["initiator"]) + require.Equal(t, "crypto_error_0x12a", ev["connection_error"]) + require.Equal(t, "foobar", ev["reason"]) +} + +func TestSentTransportParameters(t *testing.T) { + rcid := protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}) + name, ev := testEventEncoding(t, &ParametersSet{ + Initiator: InitiatorLocal, + SentBy: protocol.PerspectiveServer, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}), + InitialSourceConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + RetrySourceConnectionID: &rcid, + StatelessResetToken: &protocol.StatelessResetToken{0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff, 0x00}, + DisableActiveMigration: true, + MaxIdleTimeout: 321 * time.Millisecond, + MaxUDPPayloadSize: 1234, + AckDelayExponent: 12, + MaxAckDelay: 123 * time.Millisecond, + ActiveConnectionIDLimit: 7, + InitialMaxData: 4000, + InitialMaxStreamDataBidiLocal: 1000, + InitialMaxStreamDataBidiRemote: 2000, + InitialMaxStreamDataUni: 3000, + InitialMaxStreamsBidi: 10, + InitialMaxStreamsUni: 20, + MaxDatagramFrameSize: protocol.InvalidByteCount, + EnableResetStreamAt: true, + }) + + require.Equal(t, "transport:parameters_set", name) + require.Equal(t, "local", ev["initiator"]) + require.Equal(t, "deadc0de", ev["original_destination_connection_id"]) + require.Equal(t, "deadbeef", ev["initial_source_connection_id"]) + require.Equal(t, "decafbad", ev["retry_source_connection_id"]) + require.Equal(t, "112233445566778899aabbccddeeff00", ev["stateless_reset_token"]) + require.Equal(t, float64(321), ev["max_idle_timeout"]) + require.Equal(t, float64(1234), ev["max_udp_payload_size"]) + require.Equal(t, float64(12), ev["ack_delay_exponent"]) + require.Equal(t, float64(7), ev["active_connection_id_limit"]) + require.Equal(t, float64(4000), ev["initial_max_data"]) + require.Equal(t, float64(1000), ev["initial_max_stream_data_bidi_local"]) + require.Equal(t, float64(2000), ev["initial_max_stream_data_bidi_remote"]) + require.Equal(t, float64(3000), ev["initial_max_stream_data_uni"]) + require.Equal(t, float64(10), ev["initial_max_streams_bidi"]) + require.Equal(t, float64(20), ev["initial_max_streams_uni"]) + require.True(t, ev["reset_stream_at"].(bool)) + require.NotContains(t, ev, "preferred_address") + require.NotContains(t, ev, "max_datagram_frame_size") +} + +func TestServerTransportParametersWithoutStatelessResetToken(t *testing.T) { + name, ev := testEventEncoding(t, &ParametersSet{ + Initiator: InitiatorLocal, + SentBy: protocol.PerspectiveServer, + OriginalDestinationConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}), + ActiveConnectionIDLimit: 7, + }) + + require.Equal(t, "transport:parameters_set", name) + require.NotContains(t, ev, "stateless_reset_token") +} + +func TestTransportParametersWithoutRetrySourceConnectionID(t *testing.T) { + name, ev := testEventEncoding(t, &ParametersSet{ + Initiator: InitiatorLocal, + SentBy: protocol.PerspectiveServer, + StatelessResetToken: &protocol.StatelessResetToken{0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff, 0x00}, + }) + + require.Equal(t, "transport:parameters_set", name) + require.Equal(t, "local", ev["initiator"]) + require.NotContains(t, ev, "retry_source_connection_id") +} + +func TestTransportParametersWithPreferredAddress(t *testing.T) { + t.Run("IPv4 and IPv6", func(t *testing.T) { + testTransportParametersWithPreferredAddress(t, true, true) + }) + t.Run("IPv4 only", func(t *testing.T) { + testTransportParametersWithPreferredAddress(t, true, false) + }) + t.Run("IPv6 only", func(t *testing.T) { + testTransportParametersWithPreferredAddress(t, false, true) + }) +} + +func testTransportParametersWithPreferredAddress(t *testing.T, hasIPv4, hasIPv6 bool) { + addr4 := netip.AddrPortFrom(netip.AddrFrom4([4]byte{12, 34, 56, 78}), 123) + addr6 := netip.AddrPortFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}), 456) + preferredAddress := &PreferredAddress{ + ConnectionID: protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), + StatelessResetToken: protocol.StatelessResetToken{15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0}, + } + if hasIPv4 { + preferredAddress.IPv4 = addr4 + } + if hasIPv6 { + preferredAddress.IPv6 = addr6 + } + name, ev := testEventEncoding(t, &ParametersSet{ + Initiator: InitiatorLocal, + SentBy: protocol.PerspectiveServer, + PreferredAddress: preferredAddress, + }) + + require.Equal(t, "transport:parameters_set", name) + require.Equal(t, "local", ev["initiator"]) + require.Contains(t, ev, "preferred_address") + pa := ev["preferred_address"].(map[string]any) + if hasIPv4 { + require.Equal(t, "12.34.56.78", pa["ip_v4"]) + require.Equal(t, float64(123), pa["port_v4"]) + } else { + require.NotContains(t, pa, "ip_v4") + require.NotContains(t, pa, "port_v4") + } + if hasIPv6 { + require.Equal(t, "102:304:506:708:90a:b0c:d0e:f10", pa["ip_v6"]) + require.Equal(t, float64(456), pa["port_v6"]) + } else { + require.NotContains(t, pa, "ip_v6") + require.NotContains(t, pa, "port_v6") + } + require.Equal(t, "0807060504030201", pa["connection_id"]) + require.Equal(t, "0f0e0d0c0b0a09080706050403020100", pa["stateless_reset_token"]) +} + +func TestTransportParametersWithDatagramExtension(t *testing.T) { + name, ev := testEventEncoding(t, &ParametersSet{ + Initiator: InitiatorLocal, + SentBy: protocol.PerspectiveServer, + MaxDatagramFrameSize: 1337, + }) + + require.Equal(t, "transport:parameters_set", name) + require.Equal(t, float64(1337), ev["max_datagram_frame_size"]) +} + +func TestReceivedTransportParameters(t *testing.T) { + name, ev := testEventEncoding(t, &ParametersSet{ + Initiator: InitiatorRemote, + SentBy: protocol.PerspectiveClient, + }) + + require.Equal(t, "transport:parameters_set", name) + require.Equal(t, "remote", ev["initiator"]) + require.NotContains(t, ev, "original_destination_connection_id") +} + +func TestRestoredTransportParameters(t *testing.T) { + name, ev := testEventEncoding(t, &ParametersSet{ + Restore: true, + InitialMaxStreamDataBidiLocal: 100, + InitialMaxStreamDataBidiRemote: 200, + InitialMaxStreamDataUni: 300, + InitialMaxData: 400, + MaxIdleTimeout: 123 * time.Millisecond, + }) + + require.Equal(t, "transport:parameters_restored", name) + require.NotContains(t, ev, "initiator") + require.NotContains(t, ev, "original_destination_connection_id") + require.NotContains(t, ev, "stateless_reset_token") + require.NotContains(t, ev, "retry_source_connection_id") + require.NotContains(t, ev, "initial_source_connection_id") + require.Equal(t, float64(123), ev["max_idle_timeout"]) + require.Equal(t, float64(400), ev["initial_max_data"]) + require.Equal(t, float64(100), ev["initial_max_stream_data_bidi_local"]) + require.Equal(t, float64(200), ev["initial_max_stream_data_bidi_remote"]) + require.Equal(t, float64(300), ev["initial_max_stream_data_uni"]) +} + +func TestPacketSent(t *testing.T) { + name, ev := testEventEncoding(t, &PacketSent{ + Header: PacketHeader{ + PacketType: PacketTypeHandshake, + PacketNumber: 1337, + Version: protocol.Version1, + SrcConnectionID: protocol.ParseConnectionID([]byte{4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + }, + Raw: RawInfo{Length: 987, PayloadLength: 1337}, + Frames: []Frame{ + {Frame: &MaxStreamDataFrame{StreamID: 42, MaximumStreamData: 987}}, + {Frame: &StreamFrame{StreamID: 123, Offset: 1234, Length: 6, Fin: true}}, + }, + ECN: ECNCE, + }) + + require.Equal(t, "transport:packet_sent", name) + require.Contains(t, ev, "raw") + raw := ev["raw"].(map[string]any) + require.NotContains(t, ev, "datagram_payload_checksum") + require.Equal(t, float64(987), raw["length"]) + require.Equal(t, float64(1337), raw["payload_length"]) + require.Contains(t, ev, "header") + hdr := ev["header"].(map[string]any) + require.Equal(t, "handshake", hdr["packet_type"]) + require.Equal(t, float64(1337), hdr["packet_number"]) + require.Equal(t, "04030201", hdr["scid"]) + require.Contains(t, ev, "frames") + require.Equal(t, "CE", ev["ecn"]) + frames := ev["frames"].([]any) + require.Len(t, frames, 2) + require.Equal(t, "max_stream_data", frames[0].(map[string]any)["frame_type"]) + require.Equal(t, "stream", frames[1].(map[string]any)["frame_type"]) +} + +func TestPacketSent1RTT(t *testing.T) { + t.Run("with datagram payload checksum", func(t *testing.T) { + testPacketSent1RTT(t, 1337) + }) + + t.Run("without datagram payload checksum", func(t *testing.T) { + testPacketSent1RTT(t, 0) + }) +} + +func testPacketSent1RTT(t *testing.T, datagramPayloadChecksum DatagramPayloadChecksum) { + name, ev := testEventEncoding(t, &PacketSent{ + Header: PacketHeader{ + PacketType: PacketType1RTT, + PacketNumber: 1337, + KeyPhaseBit: KeyPhaseZero, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + }, + Raw: RawInfo{Length: 123}, + Frames: []Frame{ + {Frame: &AckFrame{AckRanges: []wire.AckRange{{Smallest: 1, Largest: 10}}}}, + {Frame: &MaxDataFrame{MaximumData: 987}}, + }, + ECN: ECNUnsupported, + DatagramPayloadChecksum: datagramPayloadChecksum, + }) + + require.Equal(t, "transport:packet_sent", name) + raw := ev["raw"].(map[string]any) + require.Equal(t, float64(123), raw["length"]) + require.NotContains(t, raw, "payload_length") + require.Contains(t, ev, "header") + require.NotContains(t, ev, "ecn") + hdr := ev["header"].(map[string]any) + require.Equal(t, "1RTT", hdr["packet_type"]) + require.Equal(t, float64(1337), hdr["packet_number"]) + require.Contains(t, ev, "frames") + frames := ev["frames"].([]any) + require.Len(t, frames, 2) + require.Equal(t, "ack", frames[0].(map[string]any)["frame_type"]) + require.Equal(t, "max_data", frames[1].(map[string]any)["frame_type"]) + if datagramPayloadChecksum != 0 { + require.Contains(t, ev, "datagram_payload_checksum") + require.Equal(t, float64(datagramPayloadChecksum), ev["datagram_payload_checksum"]) + } else { + require.NotContains(t, ev, "datagram_payload_checksum") + } +} + +func TestPacketReceived(t *testing.T) { + name, ev := testEventEncoding(t, &PacketReceived{ + Header: PacketHeader{ + PacketType: PacketTypeInitial, + PacketNumber: 1337, + Version: protocol.Version1, + SrcConnectionID: protocol.ParseConnectionID([]byte{4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Token: &Token{Raw: []byte{0xde, 0xad, 0xbe, 0xef}}, + }, + Raw: RawInfo{ + Length: 789, + PayloadLength: 1234, + }, + Frames: []Frame{ + {Frame: &MaxStreamDataFrame{StreamID: 42, MaximumStreamData: 987}}, + {Frame: &StreamFrame{StreamID: 123, Offset: 1234, Length: 6, Fin: true}}, + }, + ECN: ECT0, + DatagramPayloadChecksum: 42, + }) + + require.Equal(t, "transport:packet_received", name) + require.Contains(t, ev, "raw") + raw := ev["raw"].(map[string]any) + require.Equal(t, float64(789), raw["length"]) + require.Equal(t, float64(1234), raw["payload_length"]) + require.Equal(t, "ECT(0)", ev["ecn"]) + require.Contains(t, ev, "header") + hdr := ev["header"].(map[string]any) + require.Equal(t, "initial", hdr["packet_type"]) + require.Equal(t, float64(1337), hdr["packet_number"]) + require.Equal(t, "04030201", hdr["scid"]) + require.Contains(t, hdr, "token") + token := hdr["token"].(map[string]any) + require.Equal(t, "deadbeef", token["data"]) + require.Contains(t, ev, "frames") + require.Len(t, ev["frames"].([]any), 2) + require.Contains(t, ev, "datagram_payload_checksum") + require.Equal(t, float64(42), ev["datagram_payload_checksum"]) +} + +func TestPacketReceived1RTT(t *testing.T) { + t.Run("with datagram payload checksum", func(t *testing.T) { + testPacketReceived1RTT(t, 1337) + }) + + t.Run("without datagram payload checksum", func(t *testing.T) { + testPacketReceived1RTT(t, 0) + }) +} + +func testPacketReceived1RTT(t *testing.T, datagramPayloadChecksum DatagramPayloadChecksum) { + name, ev := testEventEncoding(t, &PacketReceived{ + Header: PacketHeader{ + PacketType: PacketType1RTT, + PacketNumber: 1337, + KeyPhaseBit: KeyPhaseZero, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + }, + Raw: RawInfo{Length: 789, PayloadLength: 1234}, + Frames: []Frame{ + {Frame: &MaxStreamDataFrame{StreamID: 42, MaximumStreamData: 987}}, + {Frame: &StreamFrame{StreamID: 123, Offset: 1234, Length: 6, Fin: true}}, + }, + ECN: ECT1, + DatagramPayloadChecksum: datagramPayloadChecksum, + }) + + require.Equal(t, "transport:packet_received", name) + require.Contains(t, ev, "raw") + raw := ev["raw"].(map[string]any) + require.Equal(t, float64(789), raw["length"]) + require.Equal(t, float64(1234), raw["payload_length"]) + require.Equal(t, "ECT(1)", ev["ecn"]) + require.Contains(t, ev, "header") + hdr := ev["header"].(map[string]any) + require.Equal(t, "1RTT", hdr["packet_type"]) + require.Equal(t, float64(1337), hdr["packet_number"]) + require.Contains(t, ev, "frames") + require.Len(t, ev["frames"].([]any), 2) + if datagramPayloadChecksum != 0 { + require.Contains(t, ev, "datagram_payload_checksum") + require.Equal(t, float64(datagramPayloadChecksum), ev["datagram_payload_checksum"]) + } else { + require.NotContains(t, ev, "datagram_payload_checksum") + } +} + +func TestPacketReceivedRetry(t *testing.T) { + name, ev := testEventEncoding(t, &PacketReceived{ + Header: PacketHeader{ + PacketType: PacketTypeRetry, + Version: protocol.Version1, + SrcConnectionID: protocol.ParseConnectionID([]byte{4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Token: &Token{Raw: []byte{0xde, 0xad, 0xbe, 0xef}}, + }, + Raw: RawInfo{Length: 123}, + }) + + require.Equal(t, "transport:packet_received", name) + require.Contains(t, ev, "raw") + raw := ev["raw"].(map[string]any) + require.Len(t, raw, 1) + require.Equal(t, float64(123), raw["length"]) + require.Contains(t, ev, "header") + header := ev["header"].(map[string]any) + require.Equal(t, "retry", header["packet_type"]) + require.NotContains(t, header, "packet_number") + require.Contains(t, header, "version") + require.Contains(t, header, "dcid") + require.Contains(t, header, "scid") + require.Contains(t, header, "token") + token := header["token"].(map[string]any) + require.Equal(t, "deadbeef", token["data"]) + require.NotContains(t, ev, "frames") +} + +func TestVersionNegotiationReceived(t *testing.T) { + name, ev := testEventEncoding(t, &VersionNegotiationReceived{ + Header: PacketHeaderVersionNegotiation{ + SrcConnectionID: ArbitraryLenConnectionID{4, 3, 2, 1}, + DestConnectionID: ArbitraryLenConnectionID{1, 2, 3, 4, 5, 6, 7, 8}, + }, + SupportedVersions: []Version{0xdeadbeef, 0xdecafbad}, + }) + + require.Equal(t, "transport:packet_received", name) + require.Contains(t, ev, "header") + require.NotContains(t, ev, "frames") + require.Contains(t, ev, "supported_versions") + require.Equal(t, []any{"deadbeef", "decafbad"}, ev["supported_versions"]) + header := ev["header"].(map[string]any) + require.Equal(t, "version_negotiation", header["packet_type"]) + require.NotContains(t, header, "packet_number") + require.NotContains(t, header, "version") + require.Equal(t, "0102030405060708", header["dcid"]) + require.Equal(t, "04030201", header["scid"]) +} + +func TestPacketBuffered(t *testing.T) { + name, ev := testEventEncoding(t, &PacketBuffered{ + Header: PacketHeader{ + PacketType: PacketTypeHandshake, + PacketNumber: protocol.InvalidPacketNumber, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + SrcConnectionID: protocol.ParseConnectionID([]byte{4, 3, 2, 1}), + }, + Raw: RawInfo{Length: 1337}, + DatagramPayloadChecksum: 42, + }) + + require.Equal(t, "transport:packet_buffered", name) + require.Contains(t, ev, "header") + require.Contains(t, ev, "raw") + require.Equal(t, float64(1337), ev["raw"].(map[string]any)["length"]) + require.Equal(t, float64(42), ev["datagram_payload_checksum"]) + require.Contains(t, ev, "trigger") + require.Equal(t, "keys_unavailable", ev["trigger"]) +} + +func TestPacketDropped(t *testing.T) { + name, ev := testEventEncoding(t, &PacketDropped{ + Header: PacketHeader{PacketType: PacketTypeRetry}, + Raw: RawInfo{Length: 1337}, + DatagramPayloadChecksum: 42, + Trigger: PacketDropPayloadDecryptError, + }) + + require.Equal(t, "transport:packet_dropped", name) + require.Contains(t, ev, "raw") + require.Equal(t, float64(1337), ev["raw"].(map[string]any)["length"]) + require.Equal(t, float64(42), ev["datagram_payload_checksum"]) + require.Contains(t, ev, "header") + require.Equal(t, "payload_decrypt_error", ev["trigger"]) +} + +func TestMetricsUpdated(t *testing.T) { + rttStats := utils.NewRTTStats() + rttStats.UpdateRTT(15*time.Millisecond, 0) + rttStats.UpdateRTT(20*time.Millisecond, 0) + rttStats.UpdateRTT(25*time.Millisecond, 0) + name, ev := testEventEncoding(t, &MetricsUpdated{ + MinRTT: rttStats.MinRTT(), + SmoothedRTT: rttStats.SmoothedRTT(), + LatestRTT: rttStats.LatestRTT(), + RTTVariance: rttStats.MeanDeviation(), + CongestionWindow: 4321, + BytesInFlight: 1234, + PacketsInFlight: 42, + }) + + require.Equal(t, "recovery:metrics_updated", name) + require.Equal(t, float64(15), ev["min_rtt"]) + require.Equal(t, float64(25), ev["latest_rtt"]) + require.Contains(t, ev, "smoothed_rtt") + require.InDelta(t, rttStats.SmoothedRTT().Milliseconds(), ev["smoothed_rtt"], float64(1)) + require.Contains(t, ev, "rtt_variance") + require.InDelta(t, rttStats.MeanDeviation().Milliseconds(), ev["rtt_variance"], float64(1)) + require.Equal(t, float64(4321), ev["congestion_window"]) + require.Equal(t, float64(1234), ev["bytes_in_flight"]) + require.Equal(t, float64(42), ev["packets_in_flight"]) +} + +func TestPacketLost(t *testing.T) { + name, ev := testEventEncoding(t, &PacketLost{ + Header: PacketHeader{PacketType: PacketTypeHandshake, PacketNumber: 42}, + Trigger: PacketLossReorderingThreshold, + }) + + require.Equal(t, "recovery:packet_lost", name) + require.Contains(t, ev, "header") + require.Equal(t, "reordering_threshold", ev["trigger"]) +} + +func TestSpuriousLoss(t *testing.T) { + name, ev := testEventEncoding(t, &SpuriousLoss{ + EncryptionLevel: protocol.Encryption1RTT, + PacketNumber: 42, + PacketReordering: 1, + TimeReordering: 1337 * time.Millisecond, + }) + + require.Equal(t, "recovery:spurious_loss", name) + require.Contains(t, ev, "packet_number") + require.Equal(t, float64(42), ev["packet_number"]) + require.Contains(t, ev, "reordering_packets") + require.Equal(t, float64(1), ev["reordering_packets"]) + require.Contains(t, ev, "reordering_time") + require.InDelta(t, 1337, ev["reordering_time"], float64(1)) +} + +func TestMTUUpdated(t *testing.T) { + name, ev := testEventEncoding(t, &MTUUpdated{ + Value: 1337, + Done: true, + }) + + require.Equal(t, "recovery:mtu_updated", name) + require.Equal(t, float64(1337), ev["mtu"]) + require.Equal(t, true, ev["done"]) +} + +func TestCongestionStateUpdated(t *testing.T) { + name, ev := testEventEncoding(t, &CongestionStateUpdated{ + State: CongestionStateCongestionAvoidance, + }) + + require.Equal(t, "recovery:congestion_state_updated", name) + require.Equal(t, "congestion_avoidance", ev["new"]) +} + +func TestPTOCountUpdated(t *testing.T) { + name, ev := testEventEncoding(t, &PTOCountUpdated{PTOCount: 42}) + + require.Equal(t, "recovery:metrics_updated", name) + require.Equal(t, float64(42), ev["pto_count"]) +} + +func TestKeyUpdatedTLS(t *testing.T) { + name, ev := testEventEncoding(t, &KeyUpdated{ + Trigger: KeyUpdateTLS, + KeyType: KeyTypeClientHandshake, + KeyPhase: 0, + }) + + require.Equal(t, "security:key_updated", name) + require.Equal(t, "client_handshake_secret", ev["key_type"]) + require.Equal(t, "tls", ev["trigger"]) + require.NotContains(t, ev, "key_phase") + require.NotContains(t, ev, "old") + require.NotContains(t, ev, "new") +} + +func TestKeyUpdatedTLS1RTT(t *testing.T) { + name, ev := testEventEncoding(t, &KeyUpdated{ + Trigger: KeyUpdateTLS, + KeyType: KeyTypeServer1RTT, + KeyPhase: 0, + }) + + require.Equal(t, "security:key_updated", name) + require.Equal(t, "server_1rtt_secret", ev["key_type"]) + require.Equal(t, "tls", ev["trigger"]) + require.Equal(t, float64(0), ev["key_phase"]) + require.NotContains(t, ev, "old") + require.NotContains(t, ev, "new") +} + +func TestKeyUpdated(t *testing.T) { + name, ev := testEventEncoding(t, &KeyUpdated{ + Trigger: KeyUpdateRemote, + KeyType: KeyTypeClient1RTT, + KeyPhase: 1337, + }) + + require.Equal(t, "security:key_updated", name) + require.Equal(t, float64(1337), ev["key_phase"]) + require.Equal(t, "remote_update", ev["trigger"]) + require.Contains(t, ev, "key_type") + require.Equal(t, "client_1rtt_secret", ev["key_type"]) +} + +func TestKeyDiscarded0RTT(t *testing.T) { + name, ev := testEventEncoding(t, &KeyDiscarded{ + KeyType: KeyTypeServer0RTT, + KeyPhase: 0, + }) + + require.Equal(t, "security:key_discarded", name) + require.Equal(t, "tls", ev["trigger"]) + require.Equal(t, "server_0rtt_secret", ev["key_type"]) +} + +func TestKeyDiscarded(t *testing.T) { + name, ev := testEventEncoding(t, &KeyDiscarded{ + KeyType: KeyTypeClient1RTT, + KeyPhase: 42, + }) + + require.Equal(t, "security:key_discarded", name) + require.Equal(t, float64(42), ev["key_phase"]) + require.NotContains(t, ev, "trigger") + require.Contains(t, ev, "key_type") + require.Equal(t, "client_1rtt_secret", ev["key_type"]) +} + +func TestLossTimerUpdated(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var buf bytes.Buffer + tr := qlogwriter.NewConnectionFileSeq( + nopWriteCloser(&buf), + true, + protocol.ParseConnectionID([]byte{1, 2, 3, 4}), + []string{EventSchema}, + ) + go tr.Run() + producer := tr.AddProducer() + + synctest.Wait() + time.Sleep(42 * time.Second) + + producer.RecordEvent(&LossTimerUpdated{ + Type: LossTimerUpdateTypeSet, + TimerType: TimerTypePTO, + EncLevel: protocol.EncryptionHandshake, + Time: time.Now().Add(1337 * time.Second), + }) + producer.Close() + + name, ev := decode(t, buf.String()) + require.Equal(t, "recovery:loss_timer_updated", name) + require.Len(t, ev, 4) + require.Equal(t, "set", ev["event_type"]) + require.Equal(t, "pto", ev["timer_type"]) + require.Equal(t, "handshake", ev["packet_number_space"]) + require.Contains(t, ev, "delta") + delta := time.Duration(ev["delta"].(float64)*1e6) * time.Nanosecond + require.Equal(t, 1337*time.Second, delta) + }) +} + +func TestLossTimerUpdatedExpired(t *testing.T) { + name, ev := testEventEncoding(t, &LossTimerUpdated{ + Type: LossTimerUpdateTypeExpired, + TimerType: TimerTypeACK, + EncLevel: protocol.Encryption1RTT, + }) + + require.Equal(t, "recovery:loss_timer_updated", name) + require.Len(t, ev, 3) + require.Equal(t, "expired", ev["event_type"]) + require.Equal(t, "ack", ev["timer_type"]) + require.Equal(t, "application_data", ev["packet_number_space"]) +} + +func TestLossTimerUpdatedCanceled(t *testing.T) { + name, ev := testEventEncoding(t, &eventLossTimerCanceled{}) + + require.Equal(t, "recovery:loss_timer_updated", name) + require.Len(t, ev, 1) + require.Equal(t, "cancelled", ev["event_type"]) +} + +func TestECNStateUpdated(t *testing.T) { + name, ev := testEventEncoding(t, &ECNStateUpdated{ + State: ECNStateUnknown, + Trigger: "", + }) + + require.Equal(t, "recovery:ecn_state_updated", name) + require.Len(t, ev, 1) + require.Equal(t, "unknown", ev["new"]) +} + +func TestECNStateUpdatedWithTrigger(t *testing.T) { + name, ev := testEventEncoding(t, &ECNStateUpdated{ + State: ECNStateFailed, + Trigger: "ACK doesn't contain ECN marks", + }) + + require.Equal(t, "recovery:ecn_state_updated", name) + require.Len(t, ev, 2) + require.Equal(t, "failed", ev["new"]) + require.Equal(t, "ACK doesn't contain ECN marks", ev["trigger"]) +} + +func TestALPNInformation(t *testing.T) { + name, ev := testEventEncoding(t, &ALPNInformation{ + ChosenALPN: "h3", + }) + + require.Equal(t, "transport:alpn_information", name) + require.Len(t, ev, 1) + require.Equal(t, "h3", ev["chosen_alpn"]) +} + +func TestDebugEvent(t *testing.T) { + t.Run("default name", func(t *testing.T) { + name, ev := testEventEncoding(t, &DebugEvent{Message: "hello world"}) + require.Equal(t, "transport:debug", name) + require.Len(t, ev, 1) + require.Equal(t, "hello world", ev["message"]) + }) + + t.Run("custom name", func(t *testing.T) { + name, ev := testEventEncoding(t, &DebugEvent{EventName: "foo", Message: "bar"}) + require.Equal(t, "transport:foo", name) + require.Len(t, ev, 1) + require.Equal(t, "bar", ev["message"]) + }) +} diff --git a/third_party/quic-go/qlog/frame.go b/third_party/quic-go/qlog/frame.go new file mode 100644 index 0000000..9a5df7d --- /dev/null +++ b/third_party/quic-go/qlog/frame.go @@ -0,0 +1,481 @@ +package qlog + +import ( + "encoding/hex" + + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +type Frame struct { + Frame any +} + +type frames []Frame + +type ( + // An AckFrame is an ACK frame. + AckFrame = wire.AckFrame + // A ConnectionCloseFrame is a CONNECTION_CLOSE frame. + ConnectionCloseFrame = wire.ConnectionCloseFrame + // A DataBlockedFrame is a DATA_BLOCKED frame. + DataBlockedFrame = wire.DataBlockedFrame + // A HandshakeDoneFrame is a HANDSHAKE_DONE frame. + HandshakeDoneFrame = wire.HandshakeDoneFrame + // A MaxDataFrame is a MAX_DATA frame. + MaxDataFrame = wire.MaxDataFrame + // A MaxStreamDataFrame is a MAX_STREAM_DATA frame. + MaxStreamDataFrame = wire.MaxStreamDataFrame + // A MaxStreamsFrame is a MAX_STREAMS_FRAME. + MaxStreamsFrame = wire.MaxStreamsFrame + // A NewConnectionIDFrame is a NEW_CONNECTION_ID frame. + NewConnectionIDFrame = wire.NewConnectionIDFrame + // A NewTokenFrame is a NEW_TOKEN frame. + NewTokenFrame = wire.NewTokenFrame + // A PathChallengeFrame is a PATH_CHALLENGE frame. + PathChallengeFrame = wire.PathChallengeFrame + // A PathResponseFrame is a PATH_RESPONSE frame. + PathResponseFrame = wire.PathResponseFrame + // A PingFrame is a PING frame. + PingFrame = wire.PingFrame + // A ResetStreamFrame is a RESET_STREAM frame. + ResetStreamFrame = wire.ResetStreamFrame + // A RetireConnectionIDFrame is a RETIRE_CONNECTION_ID frame. + RetireConnectionIDFrame = wire.RetireConnectionIDFrame + // A StopSendingFrame is a STOP_SENDING frame. + StopSendingFrame = wire.StopSendingFrame + // A StreamsBlockedFrame is a STREAMS_BLOCKED frame. + StreamsBlockedFrame = wire.StreamsBlockedFrame + // A StreamDataBlockedFrame is a STREAM_DATA_BLOCKED frame. + StreamDataBlockedFrame = wire.StreamDataBlockedFrame + // An AckFrequencyFrame is an ACK_FREQUENCY frame. + AckFrequencyFrame = wire.AckFrequencyFrame + // An ImmediateAckFrame is an IMMEDIATE_ACK frame. + ImmediateAckFrame = wire.ImmediateAckFrame +) + +type AckRange = wire.AckRange + +// A CryptoFrame is a CRYPTO frame. +type CryptoFrame struct { + Offset int64 + Length int64 +} + +// A StreamFrame is a STREAM frame. +type StreamFrame struct { + StreamID StreamID + Offset int64 + Length int64 + Fin bool +} + +// A DatagramFrame is a DATAGRAM frame. +type DatagramFrame struct { + Length int64 +} + +func (fs frames) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginArray) + for _, f := range fs { + if err := f.Encode(enc); err != nil { + return err + } + } + h.WriteToken(jsontext.EndArray) + return h.err +} + +func (f Frame) Encode(enc *jsontext.Encoder) error { + switch frame := f.Frame.(type) { + case *PingFrame: + return encodePingFrame(enc, frame) + case *AckFrame: + return encodeAckFrame(enc, frame) + case *ResetStreamFrame: + return encodeResetStreamFrame(enc, frame) + case *StopSendingFrame: + return encodeStopSendingFrame(enc, frame) + case *CryptoFrame: + return encodeCryptoFrame(enc, frame) + case *NewTokenFrame: + return encodeNewTokenFrame(enc, frame) + case *StreamFrame: + return encodeStreamFrame(enc, frame) + case *MaxDataFrame: + return encodeMaxDataFrame(enc, frame) + case *MaxStreamDataFrame: + return encodeMaxStreamDataFrame(enc, frame) + case *MaxStreamsFrame: + return encodeMaxStreamsFrame(enc, frame) + case *DataBlockedFrame: + return encodeDataBlockedFrame(enc, frame) + case *StreamDataBlockedFrame: + return encodeStreamDataBlockedFrame(enc, frame) + case *StreamsBlockedFrame: + return encodeStreamsBlockedFrame(enc, frame) + case *NewConnectionIDFrame: + return encodeNewConnectionIDFrame(enc, frame) + case *RetireConnectionIDFrame: + return encodeRetireConnectionIDFrame(enc, frame) + case *PathChallengeFrame: + return encodePathChallengeFrame(enc, frame) + case *PathResponseFrame: + return encodePathResponseFrame(enc, frame) + case *ConnectionCloseFrame: + return encodeConnectionCloseFrame(enc, frame) + case *HandshakeDoneFrame: + return encodeHandshakeDoneFrame(enc, frame) + case *DatagramFrame: + return encodeDatagramFrame(enc, frame) + case *AckFrequencyFrame: + return encodeAckFrequencyFrame(enc, frame) + case *ImmediateAckFrame: + return encodeImmediateAckFrame(enc, frame) + default: + panic("unknown frame type") + } +} + +func encodePingFrame(enc *jsontext.Encoder, _ *PingFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("ping")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type ackRanges []wire.AckRange + +func (ars ackRanges) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginArray) + for _, r := range ars { + if err := ackRange(r).encode(enc); err != nil { + return err + } + } + h.WriteToken(jsontext.EndArray) + return h.err +} + +type ackRange wire.AckRange + +func (ar ackRange) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginArray) + h.WriteToken(jsontext.Int(int64(ar.Smallest))) + if ar.Smallest != ar.Largest { + h.WriteToken(jsontext.Int(int64(ar.Largest))) + } + h.WriteToken(jsontext.EndArray) + return h.err +} + +func encodeAckFrame(enc *jsontext.Encoder, f *AckFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("ack")) + if f.DelayTime > 0 { + h.WriteToken(jsontext.String("ack_delay")) + h.WriteToken(jsontext.Float(milliseconds(f.DelayTime))) + } + h.WriteToken(jsontext.String("acked_ranges")) + if err := ackRanges(f.AckRanges).encode(enc); err != nil { + return err + } + hasECN := f.ECT0 > 0 || f.ECT1 > 0 || f.ECNCE > 0 + if hasECN { + h.WriteToken(jsontext.String("ect0")) + h.WriteToken(jsontext.Uint(f.ECT0)) + h.WriteToken(jsontext.String("ect1")) + h.WriteToken(jsontext.Uint(f.ECT1)) + h.WriteToken(jsontext.String("ce")) + h.WriteToken(jsontext.Uint(f.ECNCE)) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeResetStreamFrame(enc *jsontext.Encoder, f *ResetStreamFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + if f.ReliableSize > 0 { + h.WriteToken(jsontext.String("reset_stream_at")) + } else { + h.WriteToken(jsontext.String("reset_stream")) + } + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Int(int64(f.StreamID))) + h.WriteToken(jsontext.String("error_code")) + h.WriteToken(jsontext.Int(int64(f.ErrorCode))) + h.WriteToken(jsontext.String("final_size")) + h.WriteToken(jsontext.Int(int64(f.FinalSize))) + if f.ReliableSize > 0 { + h.WriteToken(jsontext.String("reliable_size")) + h.WriteToken(jsontext.Int(int64(f.ReliableSize))) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeStopSendingFrame(enc *jsontext.Encoder, f *StopSendingFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("stop_sending")) + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Int(int64(f.StreamID))) + h.WriteToken(jsontext.String("error_code")) + h.WriteToken(jsontext.Int(int64(f.ErrorCode))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeCryptoFrame(enc *jsontext.Encoder, f *CryptoFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("crypto")) + h.WriteToken(jsontext.String("offset")) + h.WriteToken(jsontext.Int(f.Offset)) + h.WriteToken(jsontext.String("length")) + h.WriteToken(jsontext.Int(f.Length)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeNewTokenFrame(enc *jsontext.Encoder, f *NewTokenFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("new_token")) + h.WriteToken(jsontext.String("token")) + if err := (Token{Raw: f.Token}).encode(enc); err != nil { + return err + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeStreamFrame(enc *jsontext.Encoder, f *StreamFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("stream")) + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Int(int64(f.StreamID))) + h.WriteToken(jsontext.String("offset")) + h.WriteToken(jsontext.Int(f.Offset)) + h.WriteToken(jsontext.String("length")) + h.WriteToken(jsontext.Int(f.Length)) + if f.Fin { + h.WriteToken(jsontext.String("fin")) + h.WriteToken(jsontext.True) + } + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeMaxDataFrame(enc *jsontext.Encoder, f *MaxDataFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("max_data")) + h.WriteToken(jsontext.String("maximum")) + h.WriteToken(jsontext.Int(int64(f.MaximumData))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeMaxStreamDataFrame(enc *jsontext.Encoder, f *MaxStreamDataFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("max_stream_data")) + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Int(int64(f.StreamID))) + h.WriteToken(jsontext.String("maximum")) + h.WriteToken(jsontext.Int(int64(f.MaximumStreamData))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeMaxStreamsFrame(enc *jsontext.Encoder, f *MaxStreamsFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("max_streams")) + h.WriteToken(jsontext.String("stream_type")) + h.WriteToken(jsontext.String(streamType(f.Type).String())) + h.WriteToken(jsontext.String("maximum")) + h.WriteToken(jsontext.Int(int64(f.MaxStreamNum))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeDataBlockedFrame(enc *jsontext.Encoder, f *DataBlockedFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("data_blocked")) + h.WriteToken(jsontext.String("limit")) + h.WriteToken(jsontext.Int(int64(f.MaximumData))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeStreamDataBlockedFrame(enc *jsontext.Encoder, f *StreamDataBlockedFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("stream_data_blocked")) + h.WriteToken(jsontext.String("stream_id")) + h.WriteToken(jsontext.Int(int64(f.StreamID))) + h.WriteToken(jsontext.String("limit")) + h.WriteToken(jsontext.Int(int64(f.MaximumStreamData))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeStreamsBlockedFrame(enc *jsontext.Encoder, f *StreamsBlockedFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("streams_blocked")) + h.WriteToken(jsontext.String("stream_type")) + h.WriteToken(jsontext.String(streamType(f.Type).String())) + h.WriteToken(jsontext.String("limit")) + h.WriteToken(jsontext.Int(int64(f.StreamLimit))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeNewConnectionIDFrame(enc *jsontext.Encoder, f *NewConnectionIDFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("new_connection_id")) + h.WriteToken(jsontext.String("sequence_number")) + h.WriteToken(jsontext.Uint(f.SequenceNumber)) + h.WriteToken(jsontext.String("retire_prior_to")) + h.WriteToken(jsontext.Uint(f.RetirePriorTo)) + h.WriteToken(jsontext.String("length")) + h.WriteToken(jsontext.Int(int64(f.ConnectionID.Len()))) + h.WriteToken(jsontext.String("connection_id")) + h.WriteToken(jsontext.String(f.ConnectionID.String())) + h.WriteToken(jsontext.String("stateless_reset_token")) + h.WriteToken(jsontext.String(hex.EncodeToString(f.StatelessResetToken[:]))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeRetireConnectionIDFrame(enc *jsontext.Encoder, f *RetireConnectionIDFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("retire_connection_id")) + h.WriteToken(jsontext.String("sequence_number")) + h.WriteToken(jsontext.Uint(f.SequenceNumber)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodePathChallengeFrame(enc *jsontext.Encoder, f *PathChallengeFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("path_challenge")) + h.WriteToken(jsontext.String("data")) + h.WriteToken(jsontext.String(hex.EncodeToString(f.Data[:]))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodePathResponseFrame(enc *jsontext.Encoder, f *PathResponseFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("path_response")) + h.WriteToken(jsontext.String("data")) + h.WriteToken(jsontext.String(hex.EncodeToString(f.Data[:]))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeConnectionCloseFrame(enc *jsontext.Encoder, f *ConnectionCloseFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("connection_close")) + h.WriteToken(jsontext.String("error_space")) + errorSpace := "transport" + if f.IsApplicationError { + errorSpace = "application" + } + h.WriteToken(jsontext.String(errorSpace)) + errName := transportError(f.ErrorCode).String() + if len(errName) > 0 { + h.WriteToken(jsontext.String("error_code")) + h.WriteToken(jsontext.String(errName)) + } else { + h.WriteToken(jsontext.String("error_code")) + h.WriteToken(jsontext.Uint(f.ErrorCode)) + } + h.WriteToken(jsontext.String("raw_error_code")) + h.WriteToken(jsontext.Uint(f.ErrorCode)) + h.WriteToken(jsontext.String("reason")) + h.WriteToken(jsontext.String(f.ReasonPhrase)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeHandshakeDoneFrame(enc *jsontext.Encoder, _ *HandshakeDoneFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("handshake_done")) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeDatagramFrame(enc *jsontext.Encoder, f *DatagramFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("datagram")) + h.WriteToken(jsontext.String("length")) + h.WriteToken(jsontext.Int(f.Length)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeAckFrequencyFrame(enc *jsontext.Encoder, f *AckFrequencyFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("ack_frequency")) + h.WriteToken(jsontext.String("sequence_number")) + h.WriteToken(jsontext.Uint(f.SequenceNumber)) + h.WriteToken(jsontext.String("ack_eliciting_threshold")) + h.WriteToken(jsontext.Uint(f.AckElicitingThreshold)) + h.WriteToken(jsontext.String("request_max_ack_delay")) + h.WriteToken(jsontext.Float(milliseconds(f.RequestMaxAckDelay))) + h.WriteToken(jsontext.String("reordering_threshold")) + h.WriteToken(jsontext.Int(int64(f.ReorderingThreshold))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +func encodeImmediateAckFrame(enc *jsontext.Encoder, _ *ImmediateAckFrame) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("frame_type")) + h.WriteToken(jsontext.String("immediate_ack")) + h.WriteToken(jsontext.EndObject) + return h.err +} diff --git a/third_party/quic-go/qlog/frame_test.go b/third_party/quic-go/qlog/frame_test.go new file mode 100644 index 0000000..f1e0bb6 --- /dev/null +++ b/third_party/quic-go/qlog/frame_test.go @@ -0,0 +1,421 @@ +package qlog + +import ( + "bytes" + "encoding/json" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/qlogwriter/jsontext" + + "github.com/stretchr/testify/require" +) + +func check(t *testing.T, f any, expected map[string]any) { + var buf bytes.Buffer + enc := jsontext.NewEncoder(&buf) + require.NoError(t, (Frame{Frame: f}).Encode(enc)) + data := buf.Bytes() + require.True(t, json.Valid(data)) + checkEncoding(t, data, expected) +} + +func TestPingFrame(t *testing.T) { + check(t, &PingFrame{}, map[string]any{"frame_type": "ping"}) +} + +func TestAckFrame(t *testing.T) { + tests := []struct { + name string + frame *AckFrame + expected map[string]any + }{ + { + name: "with delay and single packet range", + frame: &AckFrame{ + DelayTime: 86 * time.Millisecond, + AckRanges: []AckRange{{Smallest: 120, Largest: 120}}, + }, + expected: map[string]any{ + "frame_type": "ack", + "ack_delay": 86, + "acked_ranges": [][]float64{{120}}, + }, + }, + { + name: "without delay", + frame: &AckFrame{ + AckRanges: []AckRange{{Smallest: 120, Largest: 120}}, + }, + expected: map[string]any{ + "frame_type": "ack", + "acked_ranges": [][]float64{{120}}, + }, + }, + { + name: "with ECN counts", + frame: &AckFrame{ + AckRanges: []AckRange{{Smallest: 120, Largest: 120}}, + ECT0: 10, + ECT1: 100, + ECNCE: 1000, + }, + expected: map[string]any{ + "frame_type": "ack", + "acked_ranges": [][]float64{{120}}, + "ect0": 10, + "ect1": 100, + "ce": 1000, + }, + }, + { + name: "with multiple ranges", + frame: &AckFrame{ + DelayTime: 86 * time.Millisecond, + AckRanges: []AckRange{ + {Smallest: 5, Largest: 50}, + {Smallest: 100, Largest: 120}, + }, + }, + expected: map[string]any{ + "frame_type": "ack", + "ack_delay": 86, + "acked_ranges": [][]float64{ + {5, 50}, + {100, 120}, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + check(t, tt.frame, tt.expected) + }) + } +} + +func TestResetStreamFrame(t *testing.T) { + check(t, + &ResetStreamFrame{ + StreamID: 987, + FinalSize: 1234, + ErrorCode: 42, + }, + map[string]any{ + "frame_type": "reset_stream", + "stream_id": 987, + "error_code": 42, + "final_size": 1234, + }, + ) +} + +func TestResetStreamAtFrame(t *testing.T) { + check(t, + &ResetStreamFrame{ + StreamID: 987, + FinalSize: 1234, + ErrorCode: 42, + ReliableSize: 999, + }, + map[string]any{ + "frame_type": "reset_stream_at", + "stream_id": 987, + "error_code": 42, + "final_size": 1234, + "reliable_size": 999, + }, + ) +} + +func TestAckFrequencyFrame(t *testing.T) { + check(t, + &AckFrequencyFrame{ + SequenceNumber: 1337, + AckElicitingThreshold: 123, + RequestMaxAckDelay: 42 * time.Millisecond, + ReorderingThreshold: 1234, + }, + map[string]any{ + "frame_type": "ack_frequency", + "sequence_number": 1337, + "ack_eliciting_threshold": 123, + "request_max_ack_delay": 42, + "reordering_threshold": 1234, + }, + ) +} + +func TestImmediateAckFrame(t *testing.T) { + check(t, + &ImmediateAckFrame{}, + map[string]any{ + "frame_type": "immediate_ack", + }, + ) +} + +func TestStopSendingFrame(t *testing.T) { + check(t, + &StopSendingFrame{StreamID: 987, ErrorCode: 42}, + map[string]any{ + "frame_type": "stop_sending", + "stream_id": 987, + "error_code": 42, + }, + ) +} + +func TestCryptoFrame(t *testing.T) { + check(t, + &CryptoFrame{Offset: 1337, Length: 6}, + map[string]any{ + "frame_type": "crypto", + "offset": 1337, + "length": 6, + }, + ) +} + +func TestNewTokenFrame(t *testing.T) { + check(t, + &NewTokenFrame{Token: []byte{0xde, 0xad, 0xbe, 0xef}}, + map[string]any{ + "frame_type": "new_token", + "token": map[string]any{"data": "deadbeef"}, + }, + ) +} + +func TestStreamFrame(t *testing.T) { + tests := []struct { + name string + frame *StreamFrame + expected map[string]any + }{ + { + name: "with FIN", + frame: &StreamFrame{ + StreamID: 42, + Offset: 1337, + Fin: true, + Length: 9876, + }, + expected: map[string]any{ + "frame_type": "stream", + "stream_id": 42, + "offset": 1337, + "fin": true, + "length": 9876, + }, + }, + { + name: "without FIN", + frame: &StreamFrame{ + StreamID: 42, + Offset: 1337, + Length: 3, + }, + expected: map[string]any{ + "frame_type": "stream", + "stream_id": 42, + "offset": 1337, + "length": 3, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + check(t, tt.frame, tt.expected) + }) + } +} + +func TestMaxDataFrame(t *testing.T) { + check(t, + &MaxDataFrame{MaximumData: 1337}, + map[string]any{ + "frame_type": "max_data", + "maximum": 1337, + }, + ) +} + +func TestMaxStreamDataFrame(t *testing.T) { + check(t, + &MaxStreamDataFrame{StreamID: 1234, MaximumStreamData: 1337}, + map[string]any{ + "frame_type": "max_stream_data", + "stream_id": 1234, + "maximum": 1337, + }, + ) +} + +func TestMaxStreamsFrame(t *testing.T) { + check(t, + &MaxStreamsFrame{ + Type: protocol.StreamTypeBidi, + MaxStreamNum: 42, + }, + map[string]any{ + "frame_type": "max_streams", + "stream_type": "bidirectional", + "maximum": 42, + }, + ) +} + +func TestDataBlockedFrame(t *testing.T) { + check(t, + &DataBlockedFrame{MaximumData: 1337}, + map[string]any{ + "frame_type": "data_blocked", + "limit": 1337, + }, + ) +} + +func TestStreamDataBlockedFrame(t *testing.T) { + check(t, + &StreamDataBlockedFrame{ + StreamID: 42, + MaximumStreamData: 1337, + }, + map[string]any{ + "frame_type": "stream_data_blocked", + "stream_id": 42, + "limit": 1337, + }, + ) +} + +func TestStreamsBlockedFrame(t *testing.T) { + check(t, + &StreamsBlockedFrame{ + Type: protocol.StreamTypeUni, + StreamLimit: 123, + }, + map[string]any{ + "frame_type": "streams_blocked", + "stream_type": "unidirectional", + "limit": 123, + }, + ) +} + +func TestNewConnectionIDFrame(t *testing.T) { + check(t, + &NewConnectionIDFrame{ + SequenceNumber: 42, + RetirePriorTo: 24, + ConnectionID: protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + StatelessResetToken: protocol.StatelessResetToken{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0xa, 0xb, 0xc, 0xd, 0xe, 0xf}, + }, + map[string]any{ + "frame_type": "new_connection_id", + "sequence_number": 42, + "retire_prior_to": 24, + "length": 4, + "connection_id": "deadbeef", + "stateless_reset_token": "000102030405060708090a0b0c0d0e0f", + }, + ) +} + +func TestRetireConnectionIDFrame(t *testing.T) { + check(t, + &RetireConnectionIDFrame{SequenceNumber: 1337}, + map[string]any{ + "frame_type": "retire_connection_id", + "sequence_number": 1337, + }, + ) +} + +func TestPathChallengeFrame(t *testing.T) { + check(t, + &PathChallengeFrame{Data: [8]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0xc0, 0x01}}, + map[string]any{ + "frame_type": "path_challenge", + "data": "deadbeefcafec001", + }, + ) +} + +func TestPathResponseFrame(t *testing.T) { + check(t, + &PathResponseFrame{Data: [8]byte{0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0xc0, 0x01}}, + map[string]any{ + "frame_type": "path_response", + "data": "deadbeefcafec001", + }, + ) +} + +func TestConnectionCloseFrame(t *testing.T) { + tests := []struct { + name string + frame *ConnectionCloseFrame + expected map[string]any + }{ + { + name: "application error code", + frame: &ConnectionCloseFrame{ + IsApplicationError: true, + ErrorCode: 1337, + ReasonPhrase: "lorem ipsum", + }, + expected: map[string]any{ + "frame_type": "connection_close", + "error_space": "application", + "error_code": 1337, + "raw_error_code": 1337, + "reason": "lorem ipsum", + }, + }, + { + name: "transport error code", + frame: &ConnectionCloseFrame{ + ErrorCode: uint64(qerr.FlowControlError), + ReasonPhrase: "lorem ipsum", + }, + expected: map[string]any{ + "frame_type": "connection_close", + "error_space": "transport", + "error_code": "flow_control_error", + "raw_error_code": int(qerr.FlowControlError), + "reason": "lorem ipsum", + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + check(t, tt.frame, tt.expected) + }) + } +} + +func TestHandshakeDoneFrame(t *testing.T) { + check(t, + &HandshakeDoneFrame{}, + map[string]any{ + "frame_type": "handshake_done", + }, + ) +} + +func TestDatagramFrame(t *testing.T) { + check(t, + &DatagramFrame{Length: 1337}, + map[string]any{ + "frame_type": "datagram", + "length": 1337, + }, + ) +} diff --git a/third_party/quic-go/qlog/json_helper_test.go b/third_party/quic-go/qlog/json_helper_test.go new file mode 100644 index 0000000..648c97d --- /dev/null +++ b/third_party/quic-go/qlog/json_helper_test.go @@ -0,0 +1,42 @@ +package qlog + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +func checkEncoding(t *testing.T, data []byte, expected map[string]any) { + t.Helper() + + m := make(map[string]any) + require.NoError(t, json.Unmarshal(data, &m)) + require.Len(t, m, len(expected)) + + for key, value := range expected { + switch v := value.(type) { + case bool, string, map[string]any: + require.Equal(t, v, m[key]) + case int: + require.Equal(t, float64(v), m[key]) + case [][]float64: // used in the ACK frame + require.Contains(t, m, key) + outerSlice, ok := m[key].([]any) + require.True(t, ok) + require.Len(t, outerSlice, len(v)) + for i, innerExpected := range v { + innerSlice, ok := outerSlice[i].([]any) + require.True(t, ok) + require.Len(t, innerSlice, len(innerExpected)) + for j, expectedValue := range innerExpected { + v, ok := innerSlice[j].(float64) + require.True(t, ok) + require.Equal(t, expectedValue, v) + } + } + default: + t.Fatalf("unexpected type: %T", v) + } + } +} diff --git a/third_party/quic-go/qlog/packet_header.go b/third_party/quic-go/qlog/packet_header.go new file mode 100644 index 0000000..7504ae6 --- /dev/null +++ b/third_party/quic-go/qlog/packet_header.go @@ -0,0 +1,96 @@ +package qlog + +import ( + "encoding/hex" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +type Token struct { + Raw []byte +} + +func (t Token) encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("data")) + h.WriteToken(jsontext.String(hex.EncodeToString(t.Raw))) + h.WriteToken(jsontext.EndObject) + return h.err +} + +// PacketHeader is a QUIC packet header. +type PacketHeader struct { + PacketType PacketType + KeyPhaseBit KeyPhaseBit + PacketNumber PacketNumber + Version Version + SrcConnectionID ConnectionID + DestConnectionID ConnectionID + Token *Token +} + +func (h PacketHeader) encode(enc *jsontext.Encoder) error { + helper := encoderHelper{enc: enc} + helper.WriteToken(jsontext.BeginObject) + helper.WriteToken(jsontext.String("packet_type")) + helper.WriteToken(jsontext.String(string(h.PacketType))) + if h.PacketType != PacketTypeRetry && h.PacketType != PacketTypeVersionNegotiation && h.PacketType != "" && + h.PacketNumber != protocol.InvalidPacketNumber { + helper.WriteToken(jsontext.String("packet_number")) + helper.WriteToken(jsontext.Int(int64(h.PacketNumber))) + } + if h.Version != 0 { + helper.WriteToken(jsontext.String("version")) + helper.WriteToken(jsontext.String(version(h.Version).String())) + } + if h.PacketType != PacketType1RTT { + helper.WriteToken(jsontext.String("scil")) + helper.WriteToken(jsontext.Int(int64(h.SrcConnectionID.Len()))) + if h.SrcConnectionID.Len() > 0 { + helper.WriteToken(jsontext.String("scid")) + helper.WriteToken(jsontext.String(h.SrcConnectionID.String())) + } + } + helper.WriteToken(jsontext.String("dcil")) + helper.WriteToken(jsontext.Int(int64(h.DestConnectionID.Len()))) + if h.DestConnectionID.Len() > 0 { + helper.WriteToken(jsontext.String("dcid")) + helper.WriteToken(jsontext.String(h.DestConnectionID.String())) + } + if h.KeyPhaseBit == KeyPhaseZero || h.KeyPhaseBit == KeyPhaseOne { + helper.WriteToken(jsontext.String("key_phase_bit")) + helper.WriteToken(jsontext.String(h.KeyPhaseBit.String())) + } + if h.Token != nil { + helper.WriteToken(jsontext.String("token")) + if err := h.Token.encode(enc); err != nil { + return err + } + } + helper.WriteToken(jsontext.EndObject) + return helper.err +} + +type PacketHeaderVersionNegotiation struct { + SrcConnectionID ArbitraryLenConnectionID + DestConnectionID ArbitraryLenConnectionID +} + +func (h PacketHeaderVersionNegotiation) encode(enc *jsontext.Encoder) error { + helper := encoderHelper{enc: enc} + helper.WriteToken(jsontext.BeginObject) + helper.WriteToken(jsontext.String("packet_type")) + helper.WriteToken(jsontext.String("version_negotiation")) + helper.WriteToken(jsontext.String("scil")) + helper.WriteToken(jsontext.Int(int64(h.SrcConnectionID.Len()))) + helper.WriteToken(jsontext.String("scid")) + helper.WriteToken(jsontext.String(h.SrcConnectionID.String())) + helper.WriteToken(jsontext.String("dcil")) + helper.WriteToken(jsontext.Int(int64(h.DestConnectionID.Len()))) + helper.WriteToken(jsontext.String("dcid")) + helper.WriteToken(jsontext.String(h.DestConnectionID.String())) + helper.WriteToken(jsontext.EndObject) + return helper.err +} diff --git a/third_party/quic-go/qlog/packet_header_test.go b/third_party/quic-go/qlog/packet_header_test.go new file mode 100644 index 0000000..1ec570f --- /dev/null +++ b/third_party/quic-go/qlog/packet_header_test.go @@ -0,0 +1,134 @@ +package qlog + +import ( + "bytes" + "encoding/json" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter/jsontext" + + "github.com/stretchr/testify/require" +) + +func checkHeader(t *testing.T, hdr *PacketHeader, expected map[string]any) { + t.Helper() + + var buf bytes.Buffer + enc := jsontext.NewEncoder(&buf) + require.NoError(t, hdr.encode(enc)) + data := buf.Bytes() + require.True(t, json.Valid(data)) + checkEncoding(t, data, expected) +} + +func TestHeaderInitial(t *testing.T) { + checkHeader(t, + &PacketHeader{ + PacketType: PacketTypeInitial, + PacketNumber: 42, + Version: protocol.Version(0xdecafbad), + }, + map[string]any{ + "packet_type": "initial", + "packet_number": 42, + "dcil": 0, + "scil": 0, + "version": "decafbad", + }, + ) +} + +func TestHeaderInitialWithToken(t *testing.T) { + checkHeader(t, + &PacketHeader{ + PacketType: PacketTypeInitial, + PacketNumber: 1337, + SrcConnectionID: protocol.ParseConnectionID([]byte{0x11, 0x22, 0x33, 0x44}), + DestConnectionID: protocol.ParseConnectionID([]byte{0x55, 0x66, 0x77, 0x88}), + Version: protocol.Version(0xdecafbad), + Token: &Token{Raw: []byte{0xde, 0xad, 0xbe, 0xef}}, + }, + map[string]any{ + "packet_type": "initial", + "packet_number": 1337, + "dcil": 4, + "dcid": "55667788", + "scil": 4, + "scid": "11223344", + "version": "decafbad", + "token": map[string]any{"data": "deadbeef"}, + }, + ) +} + +func TestHeaderLongPacketNumbers(t *testing.T) { + t.Run("packet 0", func(t *testing.T) { + testHeaderPacketNumbers(t, 0) + }) + + // This is used for events where the packet number is not yet known, + // e.g. the packet_buffered event. + t.Run("no packet number", func(t *testing.T) { + testHeaderPacketNumbers(t, 1) + }) +} + +func testHeaderPacketNumbers(t *testing.T, pn protocol.PacketNumber) { + expected := map[string]any{ + "packet_type": "handshake", + "dcil": 0, + "scil": 0, + "version": "1", + } + if pn != protocol.InvalidPacketNumber { + expected["packet_number"] = int(pn) + } + checkHeader(t, + &PacketHeader{ + PacketType: PacketTypeHandshake, + PacketNumber: pn, + Version: protocol.Version1, + }, + expected, + ) +} + +func TestHeaderRetry(t *testing.T) { + checkHeader(t, + &PacketHeader{ + PacketType: PacketTypeRetry, + SrcConnectionID: protocol.ParseConnectionID([]byte{0x11, 0x22, 0x33, 0x44}), + DestConnectionID: protocol.ParseConnectionID([]byte{0x55, 0x66, 0x77, 0x88, 0x99}), + Version: protocol.Version(0xdecafbad), + Token: &Token{Raw: []byte{0xde, 0xad, 0xbe, 0xef}}, + }, + map[string]any{ + "packet_type": "retry", + "dcil": 5, + "dcid": "5566778899", + "scil": 4, + "scid": "11223344", + "token": map[string]any{"data": "deadbeef"}, + "version": "decafbad", + }, + ) +} + +func TestHeader1RTT(t *testing.T) { + checkHeader(t, + &PacketHeader{ + PacketType: PacketType1RTT, + PacketNumber: 42, + DestConnectionID: protocol.ParseConnectionID([]byte{0x55, 0x66, 0x77, 0x88}), + KeyPhaseBit: KeyPhaseZero, + }, + map[string]any{ + "packet_type": "1RTT", + "packet_number": 42, + "dcil": 4, + "dcid": "55667788", + "key_phase_bit": "0", + }, + ) +} diff --git a/third_party/quic-go/qlog/qlog_dir.go b/third_party/quic-go/qlog/qlog_dir.go new file mode 100644 index 0000000..c3a17a4 --- /dev/null +++ b/third_party/quic-go/qlog/qlog_dir.go @@ -0,0 +1,61 @@ +package qlog + +import ( + "bufio" + "context" + "fmt" + "log" + "os" + "slices" + "strings" + + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/qlogwriter" +) + +// EventSchema is the qlog event schema for QUIC +const EventSchema = "urn:ietf:params:qlog:events:quic-12" + +// DefaultConnectionTracer creates a qlog file in the qlog directory specified by the QLOGDIR environment variable. +// File names are _.sqlog. +// Returns nil if QLOGDIR is not set. +func DefaultConnectionTracer(_ context.Context, isClient bool, connID ConnectionID) qlogwriter.Trace { + return defaultConnectionTracerWithSchemas(isClient, connID, []string{EventSchema}) +} + +func DefaultConnectionTracerWithSchemas(_ context.Context, isClient bool, connID ConnectionID, eventSchemas []string) qlogwriter.Trace { + if !slices.Contains(eventSchemas, EventSchema) { + eventSchemas = append([]string{EventSchema}, eventSchemas...) + } + return defaultConnectionTracerWithSchemas(isClient, connID, eventSchemas) +} + +func defaultConnectionTracerWithSchemas(isClient bool, connID ConnectionID, eventSchemas []string) qlogwriter.Trace { + qlogDir := os.Getenv("QLOGDIR") + if qlogDir == "" { + return nil + } + if _, err := os.Stat(qlogDir); os.IsNotExist(err) { + if err := os.MkdirAll(qlogDir, 0o755); err != nil { + log.Fatalf("failed to create qlog dir %s: %v", qlogDir, err) + } + } + label := "server" + if isClient { + label = "client" + } + path := fmt.Sprintf("%s/%s_%s.sqlog", strings.TrimRight(qlogDir, "/"), connID, label) + f, err := os.Create(path) + if err != nil { + log.Printf("Failed to create qlog file %s: %s", path, err.Error()) + return nil + } + fileSeq := qlogwriter.NewConnectionFileSeq( + utils.NewBufferedWriteCloser(bufio.NewWriter(f), f), + isClient, + connID, + eventSchemas, + ) + go fileSeq.Run() + return fileSeq +} diff --git a/third_party/quic-go/qlog/qlog_dir_test.go b/third_party/quic-go/qlog/qlog_dir_test.go new file mode 100644 index 0000000..a0e2b01 --- /dev/null +++ b/third_party/quic-go/qlog/qlog_dir_test.go @@ -0,0 +1,70 @@ +package qlog + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter" + "github.com/stretchr/testify/require" +) + +func TestQLOGDIRSet(t *testing.T) { + tmpDir := t.TempDir() + + connID, _ := protocol.GenerateConnectionIDForInitial() + qlogDir := filepath.Join(tmpDir, "qlogs") + t.Setenv("QLOGDIR", qlogDir) + + t.Run("default connection tracer", func(t *testing.T) { + tracer := DefaultConnectionTracer(context.Background(), true, connID) + testQLOGDIRSet(t, qlogDir, tracer, []string{EventSchema}) + }) + + t.Run("default connection tracer with schemas", func(t *testing.T) { + tracer := DefaultConnectionTracerWithSchemas(context.Background(), true, connID, []string{"urn:ietf:params:qlog:events:foobar"}) + testQLOGDIRSet(t, qlogDir, tracer, []string{EventSchema, "urn:ietf:params:qlog:events:foobar"}) + }) +} + +func testQLOGDIRSet(t *testing.T, qlogDir string, tracer qlogwriter.Trace, expectedEventSchemas []string) { + require.NotNil(t, tracer) + + // adding and closing a producer makes the tracer close the file + recorder := tracer.AddProducer() + recorder.Close() + + _, err := os.Stat(qlogDir) + qlogDirCreated := !os.IsNotExist(err) + require.True(t, qlogDirCreated) + + entries, err := os.ReadDir(qlogDir) + require.NoError(t, err) + require.Len(t, entries, 1) + + data, err := os.ReadFile(filepath.Join(qlogDir, entries[0].Name())) + require.NoError(t, err) + + var obj map[string]any + require.NoError(t, json.Unmarshal([]byte(strings.Split(string(data), "\n")[0])[1:], &obj)) + require.Contains(t, obj, "trace") + require.IsType(t, obj["trace"], map[string]any{}) + require.Contains(t, obj["trace"], "event_schemas") + var eventSchemas []string + for _, v := range obj["trace"].(map[string]any)["event_schemas"].([]any) { + eventSchemas = append(eventSchemas, v.(string)) + } + require.Equal(t, eventSchemas, expectedEventSchemas) +} + +func TestQLOGDIRNotSet(t *testing.T) { + connID, _ := protocol.GenerateConnectionIDForInitial() + t.Setenv("QLOGDIR", "") + + tracer := DefaultConnectionTracer(context.Background(), true, connID) + require.Nil(t, tracer) +} diff --git a/third_party/quic-go/qlog/types.go b/third_party/quic-go/qlog/types.go new file mode 100644 index 0000000..1f45aa6 --- /dev/null +++ b/third_party/quic-go/qlog/types.go @@ -0,0 +1,305 @@ +package qlog + +import ( + "fmt" + "hash/crc32" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" +) + +type ( + ConnectionID = protocol.ConnectionID + ArbitraryLenConnectionID = protocol.ArbitraryLenConnectionID + Version = protocol.Version + PacketNumber = protocol.PacketNumber + EncryptionLevel = protocol.EncryptionLevel + KeyPhaseBit = protocol.KeyPhaseBit + KeyPhase = protocol.KeyPhase + StreamID = protocol.StreamID + TransportErrorCode = qerr.TransportErrorCode + ApplicationErrorCode = qerr.ApplicationErrorCode +) + +const ( + // KeyPhaseZero is key phase bit 0 + KeyPhaseZero = protocol.KeyPhaseZero + // KeyPhaseOne is key phase bit 1 + KeyPhaseOne = protocol.KeyPhaseOne +) + +// ECN represents the Explicit Congestion Notification value. +type ECN string + +const ( + // ECNUnsupported means that no ECN value was set / received + ECNUnsupported ECN = "" + // ECTNot is Not-ECT + ECTNot ECN = "Not-ECT" + // ECT0 is ECT(0) + ECT0 ECN = "ECT(0)" + // ECT1 is ECT(1) + ECT1 ECN = "ECT(1)" + // ECNCE is CE + ECNCE ECN = "CE" +) + +type Initiator string + +const ( + InitiatorLocal Initiator = "local" + InitiatorRemote Initiator = "remote" +) + +type streamType protocol.StreamType + +func (s streamType) String() string { + switch protocol.StreamType(s) { + case protocol.StreamTypeUni: + return "unidirectional" + case protocol.StreamTypeBidi: + return "bidirectional" + default: + return "unknown stream type" + } +} + +type version protocol.Version + +func (v version) String() string { + return fmt.Sprintf("%x", uint32(v)) +} + +func encLevelToPacketNumberSpace(encLevel protocol.EncryptionLevel) string { + switch encLevel { + case protocol.EncryptionInitial: + return "initial" + case protocol.EncryptionHandshake: + return "handshake" + case protocol.Encryption0RTT, protocol.Encryption1RTT: + return "application_data" + default: + return "unknown encryption level" + } +} + +// KeyType represents the type of cryptographic key used in QUIC connections. +type KeyType string + +const ( + // KeyTypeServerInitial represents the server's initial secret key. + KeyTypeServerInitial KeyType = "server_initial_secret" + // KeyTypeClientInitial represents the client's initial secret key. + KeyTypeClientInitial KeyType = "client_initial_secret" + // KeyTypeServerHandshake represents the server's handshake secret key. + KeyTypeServerHandshake KeyType = "server_handshake_secret" + // KeyTypeClientHandshake represents the client's handshake secret key. + KeyTypeClientHandshake KeyType = "client_handshake_secret" + // KeyTypeServer0RTT represents the server's 0-RTT secret key. + KeyTypeServer0RTT KeyType = "server_0rtt_secret" + // KeyTypeClient0RTT represents the client's 0-RTT secret key. + KeyTypeClient0RTT KeyType = "client_0rtt_secret" + // KeyTypeServer1RTT represents the server's 1-RTT secret key. + KeyTypeServer1RTT KeyType = "server_1rtt_secret" + // KeyTypeClient1RTT represents the client's 1-RTT secret key. + KeyTypeClient1RTT KeyType = "client_1rtt_secret" +) + +// KeyUpdateTrigger describes what caused a key update event. +type KeyUpdateTrigger string + +const ( + // KeyUpdateTLS indicates the key update was triggered by TLS. + KeyUpdateTLS KeyUpdateTrigger = "tls" + // KeyUpdateRemote indicates the key update was triggered by the remote peer. + KeyUpdateRemote KeyUpdateTrigger = "remote_update" + // KeyUpdateLocal indicates the key update was triggered locally. + KeyUpdateLocal KeyUpdateTrigger = "local_update" +) + +type transportError uint64 + +func (e transportError) String() string { + switch qerr.TransportErrorCode(e) { + case qerr.NoError: + return "no_error" + case qerr.InternalError: + return "internal_error" + case qerr.ConnectionRefused: + return "connection_refused" + case qerr.FlowControlError: + return "flow_control_error" + case qerr.StreamLimitError: + return "stream_limit_error" + case qerr.StreamStateError: + return "stream_state_error" + case qerr.FinalSizeError: + return "final_size_error" + case qerr.FrameEncodingError: + return "frame_encoding_error" + case qerr.TransportParameterError: + return "transport_parameter_error" + case qerr.ConnectionIDLimitError: + return "connection_id_limit_error" + case qerr.ProtocolViolation: + return "protocol_violation" + case qerr.InvalidToken: + return "invalid_token" + case qerr.ApplicationErrorErrorCode: + return "application_error" + case qerr.CryptoBufferExceeded: + return "crypto_buffer_exceeded" + case qerr.KeyUpdateError: + return "key_update_error" + case qerr.AEADLimitReached: + return "aead_limit_reached" + case qerr.NoViablePathError: + return "no_viable_path" + default: + return "" + } +} + +type PacketType string + +const ( + // PacketTypeInitial represents an Initial packet + PacketTypeInitial PacketType = "initial" + // PacketTypeHandshake represents a Handshake packet + PacketTypeHandshake PacketType = "handshake" + // PacketTypeRetry represents a Retry packet + PacketTypeRetry PacketType = "retry" + // PacketType0RTT represents a 0-RTT packet + PacketType0RTT PacketType = "0RTT" + // PacketTypeVersionNegotiation represents a Version Negotiation packet + PacketTypeVersionNegotiation PacketType = "version_negotiation" + // PacketTypeStatelessReset represents a Stateless Reset packet + PacketTypeStatelessReset PacketType = "stateless_reset" + // PacketType1RTT represents a 1-RTT packet + PacketType1RTT PacketType = "1RTT" + // // PacketTypeNotDetermined represents a packet type that could not be determined + // PacketTypeNotDetermined packetType = "" +) + +func EncryptionLevelToPacketType(l EncryptionLevel) PacketType { + switch l { + case protocol.EncryptionInitial: + return PacketTypeInitial + case protocol.EncryptionHandshake: + return PacketTypeHandshake + case protocol.Encryption0RTT: + return PacketType0RTT + case protocol.Encryption1RTT: + return PacketType1RTT + default: + panic(fmt.Sprintf("unknown encryption level: %d", l)) + } +} + +type PacketLossReason string + +const ( + // PacketLossReorderingThreshold is used when a packet is declared lost due to reordering threshold + PacketLossReorderingThreshold PacketLossReason = "reordering_threshold" + // PacketLossTimeThreshold is used when a packet is declared lost due to time threshold + PacketLossTimeThreshold PacketLossReason = "time_threshold" +) + +type PacketDropReason string + +const ( + // PacketDropKeyUnavailable is used when a packet is dropped because keys are unavailable + PacketDropKeyUnavailable PacketDropReason = "key_unavailable" + // PacketDropUnknownConnectionID is used when a packet is dropped because the connection ID is unknown + PacketDropUnknownConnectionID PacketDropReason = "unknown_connection_id" + // PacketDropHeaderParseError is used when a packet is dropped because header parsing failed + PacketDropHeaderParseError PacketDropReason = "header_parse_error" + // PacketDropPayloadDecryptError is used when a packet is dropped because decrypting the payload failed + PacketDropPayloadDecryptError PacketDropReason = "payload_decrypt_error" + // PacketDropProtocolViolation is used when a packet is dropped due to a protocol violation + PacketDropProtocolViolation PacketDropReason = "protocol_violation" + // PacketDropDOSPrevention is used when a packet is dropped to mitigate a DoS attack + PacketDropDOSPrevention PacketDropReason = "dos_prevention" + // PacketDropUnsupportedVersion is used when a packet is dropped because the version is not supported + PacketDropUnsupportedVersion PacketDropReason = "unsupported_version" + // PacketDropUnexpectedPacket is used when an unexpected packet is received + PacketDropUnexpectedPacket PacketDropReason = "unexpected_packet" + // PacketDropUnexpectedSourceConnectionID is used when a packet with an unexpected source connection ID is received + PacketDropUnexpectedSourceConnectionID PacketDropReason = "unexpected_source_connection_id" + // PacketDropUnexpectedVersion is used when a packet with an unexpected version is received + PacketDropUnexpectedVersion PacketDropReason = "unexpected_version" + // PacketDropDuplicate is used when a duplicate packet is received + PacketDropDuplicate PacketDropReason = "duplicate" +) + +type LossTimerUpdateType string + +const ( + LossTimerUpdateTypeSet LossTimerUpdateType = "set" + LossTimerUpdateTypeExpired LossTimerUpdateType = "expired" + LossTimerUpdateTypeCancelled LossTimerUpdateType = "cancelled" +) + +type TimerType string + +const ( + // TimerTypeACK represents an ACK timer + TimerTypeACK TimerType = "ack" + // TimerTypePTO represents a PTO (Probe Timeout) timer + TimerTypePTO TimerType = "pto" + // TimerTypePathProbe represents a path probe timer + TimerTypePathProbe TimerType = "path_probe" +) + +type CongestionState string + +const ( + // CongestionStateSlowStart is the slow start phase of Reno / Cubic + CongestionStateSlowStart CongestionState = "slow_start" + // CongestionStateCongestionAvoidance is the congestion avoidance phase of Reno / Cubic + CongestionStateCongestionAvoidance CongestionState = "congestion_avoidance" + // CongestionStateRecovery is the recovery phase of Reno / Cubic + CongestionStateRecovery CongestionState = "recovery" + // CongestionStateApplicationLimited means that the congestion controller is application limited + CongestionStateApplicationLimited CongestionState = "application_limited" +) + +func (s CongestionState) String() string { + return string(s) +} + +// ECNState is the state of the ECN state machine (see Appendix A.4 of RFC 9000) +type ECNState string + +const ( + // ECNStateTesting is the testing state + ECNStateTesting ECNState = "testing" + // ECNStateUnknown is the unknown state + ECNStateUnknown ECNState = "unknown" + // ECNStateFailed is the failed state + ECNStateFailed ECNState = "failed" + // ECNStateCapable is the capable state + ECNStateCapable ECNState = "capable" +) + +type ConnectionCloseTrigger string + +const ( + // IdleTimeout indicates the connection was closed due to idle timeout + ConnectionCloseTriggerIdleTimeout ConnectionCloseTrigger = "idle_timeout" + // Application indicates the connection was closed by the application + ConnectionCloseTriggerApplication ConnectionCloseTrigger = "application" + // VersionMismatch indicates the connection was closed due to a QUIC version mismatch + ConnectionCloseTriggerVersionMismatch ConnectionCloseTrigger = "version_mismatch" + // StatelessReset indicates the connection was closed due to receiving a stateless reset from the peer + ConnectionCloseTriggerStatelessReset ConnectionCloseTrigger = "stateless_reset" +) + +// DatagramPayloadChecksum is the CRC32c checksum of a UDP datagram payload. +// ponytail: zero means absent; use an optional value if zero-valued checksums need to be logged. +type DatagramPayloadChecksum uint32 + +// CalculateDatagramPayloadChecksum computes the checksum of a UDP datagram payload. +func CalculateDatagramPayloadChecksum(payload []byte) DatagramPayloadChecksum { + return DatagramPayloadChecksum(crc32.Checksum(payload, crc32.MakeTable(crc32.Castagnoli))) +} diff --git a/third_party/quic-go/qlog/types_test.go b/third_party/quic-go/qlog/types_test.go new file mode 100644 index 0000000..1e25ce9 --- /dev/null +++ b/third_party/quic-go/qlog/types_test.go @@ -0,0 +1,20 @@ +package qlog + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func TestEncryptionLevelToPacketType(t *testing.T) { + require.Equal(t, "initial", string(EncryptionLevelToPacketType(protocol.EncryptionInitial))) + require.Equal(t, "handshake", string(EncryptionLevelToPacketType(protocol.EncryptionHandshake))) + require.Equal(t, "0RTT", string(EncryptionLevelToPacketType(protocol.Encryption0RTT))) + require.Equal(t, "1RTT", string(EncryptionLevelToPacketType(protocol.Encryption1RTT))) +} + +func TestCalculateDatagramPayloadChecksum(t *testing.T) { + require.Equal(t, DatagramPayloadChecksum(0xe3069283), CalculateDatagramPayloadChecksum([]byte("123456789"))) +} diff --git a/third_party/quic-go/qlogwriter/jsontext/encoder.go b/third_party/quic-go/qlogwriter/jsontext/encoder.go new file mode 100644 index 0000000..4f715bb --- /dev/null +++ b/third_party/quic-go/qlogwriter/jsontext/encoder.go @@ -0,0 +1,324 @@ +// Package jsontext provides a fast JSON encoder providing only the necessary features +// for qlog encoding. No efforts are made to add any features beyond qlog's requirements. +// +// The API aims to be compatible with the standard library's encoding/json/jsontext package. +package jsontext + +import ( + "fmt" + "io" + "strconv" + "unsafe" +) + +type kind uint8 + +const ( + kindString kind = iota + kindInt + kindUint + kindFloat + kindBool + kindNull + kindObjectStart + kindObjectEnd + kindArrayStart + kindArrayEnd +) + +// Token represents a JSON token. +type Token struct { + kind kind + str string + i64 int64 + u64 uint64 + f64 float64 + b bool +} + +// String creates a string token. +func String(s string) Token { + return Token{kind: kindString, str: s} +} + +// Int creates an int token. +func Int(i int64) Token { + return Token{kind: kindInt, i64: i} +} + +// Uint creates a uint token. +func Uint(u uint64) Token { + return Token{kind: kindUint, u64: u} +} + +// Float creates a float token. +func Float(f float64) Token { + return Token{kind: kindFloat, f64: f} +} + +// Bool creates a bool token. +func Bool(b bool) Token { + return Token{kind: kindBool, b: b} +} + +// Null is a null token. +var Null Token = Token{kind: kindNull} + +// BeginObject is the begin object token. +var BeginObject Token = Token{kind: kindObjectStart} + +// EndObject is the end object token. +var EndObject Token = Token{kind: kindObjectEnd} + +// BeginArray is the begin array token. +var BeginArray Token = Token{kind: kindArrayStart} + +// EndArray is the end array token. +var EndArray Token = Token{kind: kindArrayEnd} + +// True is a true token. +var True Token = Bool(true) + +// False is a false token. +var False Token = Bool(false) + +var hexDigits = [16]byte{'0', '1', '2', '3', '4', '5', '6', '7', '8', '9', 'a', 'b', 'c', 'd', 'e', 'f'} + +var ( + commaByte = []byte(",") + quoteByte = []byte(`"`) + colonByte = []byte(":") + trueByte = []byte("true") + falseByte = []byte("false") + nullByte = []byte("null") + openObjectByte = []byte("{") + closeObjectByte = []byte("}") + openArrayByte = []byte("[") + closeArrayByte = []byte("]") + newlineByte = []byte("\n") + escapeQuote = []byte(`\"`) + escapeBackslash = []byte(`\\`) + escapeBackspace = []byte(`\b`) + escapeFormfeed = []byte(`\f`) + escapeNewline = []byte(`\n`) + escapeCarriage = []byte(`\r`) + escapeTab = []byte(`\t`) + escapeUnicode = []byte(`\u00`) +) + +type context struct { + isObject bool + needsComma bool + expectKey bool +} + +// Encoder encodes JSON to an io.Writer. +type Encoder struct { + w io.Writer + buf [64]byte // scratch buffer for number formatting + stack []context +} + +// NewEncoder creates a new Encoder. +func NewEncoder(w io.Writer) *Encoder { + stack := make([]context, 0, 8) + stack = append(stack, context{isObject: false, needsComma: false, expectKey: false}) + return &Encoder{ + w: w, + stack: stack, + } +} + +// WriteToken writes a token to the encoder. +func (e *Encoder) WriteToken(t Token) error { + if len(e.stack) == 0 { + return fmt.Errorf("empty stack") + } + curr := &e.stack[len(e.stack)-1] + isClosing := t.kind == kindObjectEnd || t.kind == kindArrayEnd + if !isClosing && curr.needsComma { + if _, err := e.w.Write(commaByte); err != nil { + return err + } + curr.needsComma = false + } + var err error + switch t.kind { + case kindString: + data := stringToBytes(t.str) + needsEscape := false + for _, b := range data { + if b == '"' || b == '\\' || b < 0x20 { + needsEscape = true + break + } + } + if !needsEscape { + if _, err = e.w.Write(quoteByte); err != nil { + return err + } + if _, err = e.w.Write(data); err != nil { + return err + } + if _, err = e.w.Write(quoteByte); err != nil { + return err + } + } else { + if _, err = e.w.Write(quoteByte); err != nil { + return err + } + for i := 0; i < len(t.str); i++ { + c := t.str[i] + switch c { + case '"': + if _, err = e.w.Write(escapeQuote); err != nil { + return err + } + case '\\': + if _, err = e.w.Write(escapeBackslash); err != nil { + return err + } + case '\b': + if _, err = e.w.Write(escapeBackspace); err != nil { + return err + } + case '\f': + if _, err = e.w.Write(escapeFormfeed); err != nil { + return err + } + case '\n': + if _, err = e.w.Write(escapeNewline); err != nil { + return err + } + case '\r': + if _, err = e.w.Write(escapeCarriage); err != nil { + return err + } + case '\t': + if _, err = e.w.Write(escapeTab); err != nil { + return err + } + default: + if c < 0x20 { + if _, err = e.w.Write(escapeUnicode); err != nil { + return err + } + if _, err = e.w.Write([]byte{hexDigits[c>>4], hexDigits[c&0xf]}); err != nil { + return err + } + } else { + if _, err = e.w.Write([]byte{c}); err != nil { + return err + } + } + } + } + if _, err = e.w.Write(quoteByte); err != nil { + return err + } + } + if curr.isObject { + if curr.expectKey { + // key + if _, err = e.w.Write(colonByte); err != nil { + return err + } + curr.expectKey = false + return nil // do not call afterValue for keys + } else { + // value + e.afterValue() + } + } else { + e.afterValue() + } + case kindInt: + b := strconv.AppendInt(e.buf[:0], t.i64, 10) + if _, err = e.w.Write(b); err != nil { + return err + } + e.afterValue() + case kindUint: + b := strconv.AppendUint(e.buf[:0], t.u64, 10) + if _, err = e.w.Write(b); err != nil { + return err + } + e.afterValue() + case kindFloat: + b := strconv.AppendFloat(e.buf[:0], t.f64, 'g', -1, 64) + if _, err = e.w.Write(b); err != nil { + return err + } + e.afterValue() + case kindBool: + if t.b { + if _, err = e.w.Write(trueByte); err != nil { + return err + } + } else { + if _, err = e.w.Write(falseByte); err != nil { + return err + } + } + e.afterValue() + case kindNull: + if _, err = e.w.Write(nullByte); err != nil { + return err + } + e.afterValue() + case kindObjectStart: + if _, err = e.w.Write(openObjectByte); err != nil { + return err + } + e.stack = append(e.stack, context{isObject: true, needsComma: false, expectKey: true}) + return nil + case kindObjectEnd: + if _, err = e.w.Write(closeObjectByte); err != nil { + return err + } + e.stack = e.stack[:len(e.stack)-1] + e.afterValue() + if len(e.stack) == 1 { + if _, err = e.w.Write(newlineByte); err != nil { + return err + } + } + return nil + case kindArrayStart: + if _, err = e.w.Write(openArrayByte); err != nil { + return err + } + e.stack = append(e.stack, context{isObject: false, needsComma: false, expectKey: false}) + return nil + case kindArrayEnd: + if _, err = e.w.Write(closeArrayByte); err != nil { + return err + } + e.stack = e.stack[:len(e.stack)-1] + e.afterValue() + if len(e.stack) == 1 { + if _, err = e.w.Write(newlineByte); err != nil { + return err + } + } + return nil + default: + return fmt.Errorf("unknown token kind") + } + return err +} + +// afterValue updates the state after encoding a value +func (e *Encoder) afterValue() { + if len(e.stack) > 1 { + curr := &e.stack[len(e.stack)-1] + curr.needsComma = true + if curr.isObject { + curr.expectKey = true + } + } +} + +func stringToBytes(s string) []byte { + return unsafe.Slice(unsafe.StringData(s), len(s)) +} diff --git a/third_party/quic-go/qlogwriter/jsontext/encoder_test.go b/third_party/quic-go/qlogwriter/jsontext/encoder_test.go new file mode 100644 index 0000000..54dbb04 --- /dev/null +++ b/third_party/quic-go/qlogwriter/jsontext/encoder_test.go @@ -0,0 +1,405 @@ +package jsontext_test + +import ( + "bytes" + "encoding/json" + "testing" + + "github.com/apernet/quic-go/qlogwriter/jsontext" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestEncoderSimpleObject(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginObject) + enc.WriteToken(jsontext.String("foo")) + enc.WriteToken(jsontext.String("bar")) + enc.WriteToken(jsontext.String("foo2")) + enc.WriteToken(jsontext.String("bar2")) + enc.WriteToken(jsontext.EndObject) + output := buf.String() + + var got map[string]string + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, map[string]string{"foo": "bar", "foo2": "bar2"}, got) +} + +func TestEncoderArrayInts(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginArray) + enc.WriteToken(jsontext.Int(1)) + enc.WriteToken(jsontext.Int(2)) + enc.WriteToken(jsontext.Int(3)) + enc.WriteToken(jsontext.EndArray) + output := buf.String() + + var got []int + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, []int{1, 2, 3}, got) +} + +func TestEncoderArrayStrings(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginArray) + enc.WriteToken(jsontext.String("one")) + enc.WriteToken(jsontext.String("two")) + enc.WriteToken(jsontext.EndArray) + output := buf.String() + + var got []string + err := json.Unmarshal([]byte(output), &got) + require.NoError(t, err) + require.Equal(t, []string{"one", "two"}, got) +} + +func TestEncoderNestedObject(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginObject) + enc.WriteToken(jsontext.String("outer")) + enc.WriteToken(jsontext.BeginObject) + enc.WriteToken(jsontext.String("inner")) + enc.WriteToken(jsontext.String("value")) + enc.WriteToken(jsontext.EndObject) + enc.WriteToken(jsontext.EndObject) + output := buf.String() + + var got map[string]map[string]string + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, map[string]map[string]string{"outer": {"inner": "value"}}, got) +} + +func TestEncoderNumbersAndBool(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginObject) + enc.WriteToken(jsontext.String("int")) + enc.WriteToken(jsontext.Int(42)) + enc.WriteToken(jsontext.String("uint")) + enc.WriteToken(jsontext.Uint(100)) + enc.WriteToken(jsontext.String("float")) + enc.WriteToken(jsontext.Float(3.14)) + enc.WriteToken(jsontext.String("true")) + enc.WriteToken(jsontext.True) + enc.WriteToken(jsontext.String("false")) + enc.WriteToken(jsontext.False) + enc.WriteToken(jsontext.String("nullv")) + enc.WriteToken(jsontext.Null) + enc.WriteToken(jsontext.EndObject) + output := buf.String() + + var got map[string]any + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, map[string]any{ + "int": float64(42), // json.Unmarshal decodes numbers as float64 + "uint": float64(100), + "float": 3.14, + "true": true, + "false": false, + "nullv": nil, + }, got) +} + +func TestEncoderEmptyObject(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginObject) + enc.WriteToken(jsontext.EndObject) + output := buf.String() + + var got map[string]any + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, map[string]any{}, got) +} + +func TestEncoderEmptyArray(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginArray) + enc.WriteToken(jsontext.EndArray) + output := buf.String() + + var got []any + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, []any{}, got) +} + +func TestEncoderArrayWithNulls(t *testing.T) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginArray) + enc.WriteToken(jsontext.Null) + enc.WriteToken(jsontext.String("x")) + enc.WriteToken(jsontext.Null) + enc.WriteToken(jsontext.EndArray) + output := buf.String() + + var got []any + require.NoError(t, json.Unmarshal([]byte(output), &got)) + require.Equal(t, []any{nil, "x", nil}, got) +} + +func TestEncoderEscapedStrings(t *testing.T) { + t.Run("no escapes", func(t *testing.T) { + testEncoderEscapedStrings(t, "simplekey", "simplevalue") + }) + + t.Run("basic escapes", func(t *testing.T) { + key := `key"\/` + value := `value"\/` + testEncoderEscapedStrings(t, key, value) + }) + + t.Run("control characters", func(t *testing.T) { + key := "key\b\f\n\r\t" + value := "value\b\f\n\r\t" + testEncoderEscapedStrings(t, key, value) + }) + + t.Run("unicode low", func(t *testing.T) { + key := "key\u0007\u001f" + value := "value\u0007\u001f" + testEncoderEscapedStrings(t, key, value) + }) + + t.Run("mixed all", func(t *testing.T) { + key := `key"\\\/\b\f\n\r\t\u0007\u001f` + value := `value"\\\/\b\f\n\r\t\u0007\u001f` + testEncoderEscapedStrings(t, key, value) + }) +} + +func testEncoderEscapedStrings(t *testing.T, key, value string) { + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + enc.WriteToken(jsontext.BeginObject) + enc.WriteToken(jsontext.String(key)) + enc.WriteToken(jsontext.String(value)) + enc.WriteToken(jsontext.EndObject) + output := buf.String() + + var got map[string]string + err := json.Unmarshal([]byte(output), &got) + require.NoError(t, err) + expected := map[string]string{key: value} + require.Equal(t, expected, got) +} + +func encodeValue(t testing.TB, enc *jsontext.Encoder, v any) (isSupported bool) { + t.Helper() + + switch val := v.(type) { + case map[string]any: + require.NoError(t, enc.WriteToken(jsontext.BeginObject)) + for k, vv := range val { + require.NoError(t, enc.WriteToken(jsontext.String(k))) + if !encodeValue(t, enc, vv) { + return false + } + } + require.NoError(t, enc.WriteToken(jsontext.EndObject)) + return true + case []any: + require.NoError(t, enc.WriteToken(jsontext.BeginArray)) + for _, vv := range val { + if !encodeValue(t, enc, vv) { + return false // Propagate unsupported if any nested value fails + } + } + require.NoError(t, enc.WriteToken(jsontext.EndArray)) + return true + case string: + require.NoError(t, enc.WriteToken(jsontext.String(val))) + return true + case int64: + require.NoError(t, enc.WriteToken(jsontext.Int(val))) + return true + case uint64: + require.NoError(t, enc.WriteToken(jsontext.Uint(val))) + return true + case float64: + require.NoError(t, enc.WriteToken(jsontext.Float(val))) + return true + case bool: + require.NoError(t, enc.WriteToken(jsontext.Bool(val))) + return true + case nil: + require.NoError(t, enc.WriteToken(jsontext.Null)) + return true + default: + return false + } +} + +type errorWriter struct { + N int +} + +func (w *errorWriter) Write(p []byte) (int, error) { + n := min(len(p), w.N) + w.N -= n + if w.N <= 0 { + return n, assert.AnError + } + return n, nil +} + +func TestEncoderComprehensive(t *testing.T) { + // encodes an object with all token types and nested structures + encode := func(enc *jsontext.Encoder) error { + if err := enc.WriteToken(jsontext.BeginObject); err != nil { + return err + } + + if err := enc.WriteToken(jsontext.String("simple")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String("value")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String("escaped")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String(`"quoted\"string"`)); err != nil { + return err + } + + if err := enc.WriteToken(jsontext.String("int")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.Int(-42)); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String("uint")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.Uint(100)); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String("float")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.Float(3.14)); err != nil { + return err + } + + if err := enc.WriteToken(jsontext.String("true")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.True); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String("false")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.False); err != nil { + return err + } + + if err := enc.WriteToken(jsontext.String("array")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.BeginArray); err != nil { + return err + } + if err := enc.WriteToken(jsontext.String("item1")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.Int(1)); err != nil { + return err + } + if err := enc.WriteToken(jsontext.EndArray); err != nil { + return err + } + + if err := enc.WriteToken(jsontext.String("nested")); err != nil { + return err + } + if err := enc.WriteToken(jsontext.BeginObject); err != nil { + return err + } + if err := enc.WriteToken(jsontext.EndObject); err != nil { + return err + } + + if err := enc.WriteToken(jsontext.EndObject); err != nil { + return err + } + return nil + } + + buf := bytes.NewBuffer(nil) + enc := jsontext.NewEncoder(buf) + require.NoError(t, encode(enc)) + + for i := range buf.Len() { + enc := jsontext.NewEncoder(&errorWriter{N: i}) + require.ErrorIs(t, encode(enc), assert.AnError) + } +} + +func FuzzEncoder(f *testing.F) { + examples := []string{ + `{"hello": "world"}`, + `{"foo": 123, "bar": [1, 2, 3]}`, + `{"nested": {"a": 1, "b": [true, false, "foobar", null]}}`, + `[{"x": 1}, {"y": "foo"}]`, + `["foo", "bar"]`, + `["a", {"b": [1, 2, {"c": "d"}]}, 3]`, + `{"emptyObj": {}, "emptyArr": []}`, + `{"mixed": [1, "two", {"three": 3}]}`, + `[null]`, + } + for _, tc := range examples { + // first test that + // 1. it's valid JSON + d := json.NewDecoder(bytes.NewReader([]byte(tc))) + var expected any + require.NoError(f, d.Decode(&expected), "corpus entry `%s` is not valid JSON", tc) + // 2. the jsontext encoder can handle + enc := jsontext.NewEncoder(&bytes.Buffer{}) + require.True(f, encodeValue(f, enc, expected), "expected `%s` to be supported", tc) + + f.Add([]byte(tc)) + } + + var stdlibBuf, ourBuf bytes.Buffer + + f.Fuzz(func(t *testing.T, b []byte) { + stdlibBuf.Truncate(0) + ourBuf.Truncate(0) + stdlibBuf.Grow(len(b)) + ourBuf.Grow(len(b)) + + d := json.NewDecoder(bytes.NewReader(b)) + var expected any + if err := d.Decode(&expected); err != nil { + return // invalid JSON + } + + // only attempt to handle inputs that the standard library can handle + stdlibEnc := json.NewEncoder(&stdlibBuf) + require.NoError(t, stdlibEnc.Encode(expected)) + if !json.Valid(stdlibBuf.Bytes()) { + return + } + + // then encode using the jsontext encoder + enc := jsontext.NewEncoder(&ourBuf) + if isSupported := encodeValue(t, enc, expected); !isSupported { + return + } + + output := ourBuf.Bytes() + require.Truef(t, json.Valid(output), "produced invalid JSON: %s", output) + + var got any + require.NoError(t, json.Unmarshal(output, &got)) + require.JSONEq(t, ourBuf.String(), stdlibBuf.String()) + }) +} diff --git a/third_party/quic-go/qlogwriter/trace.go b/third_party/quic-go/qlogwriter/trace.go new file mode 100644 index 0000000..20edf4b --- /dev/null +++ b/third_party/quic-go/qlogwriter/trace.go @@ -0,0 +1,124 @@ +package qlogwriter + +import ( + "runtime/debug" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +type ConnectionID = protocol.ConnectionID + +// Setting of this only works when quic-go is used as a library. +// When building a binary from this repository, the version can be set using the following go build flag: +// -ldflags="-X github.com/apernet/quic-go/qlogwriter.quicGoVersion=foobar" +var quicGoVersion = "(devel)" + +func init() { + if quicGoVersion != "(devel)" { // variable set by ldflags + return + } + info, ok := debug.ReadBuildInfo() + if !ok { // no build info available. This happens when quic-go is not used as a library. + return + } + for _, d := range info.Deps { + if d.Path == "github.com/apernet/quic-go" { + quicGoVersion = d.Version + if d.Replace != nil { + if len(d.Replace.Version) > 0 { + quicGoVersion = d.Version + } else { + quicGoVersion += " (replaced)" + } + } + break + } + } +} + +type encoderHelper struct { + enc *jsontext.Encoder + err error +} + +func (h *encoderHelper) WriteToken(t jsontext.Token) { + if h.err != nil { + return + } + h.err = h.enc.WriteToken(t) +} + +type traceHeader struct { + VantagePointType string + GroupID *ConnectionID + ReferenceTime time.Time + EventSchemas []string +} + +func (l traceHeader) Encode(enc *jsontext.Encoder) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("file_schema")) + h.WriteToken(jsontext.String("urn:ietf:params:qlog:file:sequential")) + h.WriteToken(jsontext.String("serialization_format")) + h.WriteToken(jsontext.String("application/qlog+json-seq")) + h.WriteToken(jsontext.String("title")) + h.WriteToken(jsontext.String("quic-go qlog")) + h.WriteToken(jsontext.String("code_version")) + h.WriteToken(jsontext.String(quicGoVersion)) + + h.WriteToken(jsontext.String("trace")) + // trace + h.WriteToken(jsontext.BeginObject) + if len(l.EventSchemas) > 0 { + h.WriteToken(jsontext.String("event_schemas")) + h.WriteToken(jsontext.BeginArray) + for _, schema := range l.EventSchemas { + h.WriteToken(jsontext.String(schema)) + } + h.WriteToken(jsontext.EndArray) + } + + h.WriteToken(jsontext.String("vantage_point")) + // -- vantage_point + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("type")) + h.WriteToken(jsontext.String(l.VantagePointType)) + // -- end vantage_point + h.WriteToken(jsontext.EndObject) + + h.WriteToken(jsontext.String("common_fields")) + // -- common_fields + h.WriteToken(jsontext.BeginObject) + if l.GroupID != nil { + h.WriteToken(jsontext.String("group_id")) + h.WriteToken(jsontext.String(l.GroupID.String())) + } + h.WriteToken(jsontext.String("reference_time")) + // ---- reference_time + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("clock_type")) + h.WriteToken(jsontext.String("monotonic")) + h.WriteToken(jsontext.String("epoch")) + h.WriteToken(jsontext.String("unknown")) + h.WriteToken(jsontext.String("wall_clock_time")) + h.WriteToken(jsontext.String(l.ReferenceTime.Format(time.RFC3339Nano))) + // ---- end reference_time + h.WriteToken(jsontext.EndObject) + // -- end common_fields + h.WriteToken(jsontext.EndObject) + // end trace + h.WriteToken(jsontext.EndObject) + + // The following fields are not required by the qlog draft anymore, + // but qvis still requires them to be present. + h.WriteToken(jsontext.String("qlog_format")) + h.WriteToken(jsontext.String("JSON-SEQ")) + h.WriteToken(jsontext.String("qlog_version")) + h.WriteToken(jsontext.String("0.3")) + + h.WriteToken(jsontext.EndObject) + return h.err +} diff --git a/third_party/quic-go/qlogwriter/trace_test.go b/third_party/quic-go/qlogwriter/trace_test.go new file mode 100644 index 0000000..e87253c --- /dev/null +++ b/third_party/quic-go/qlogwriter/trace_test.go @@ -0,0 +1,115 @@ +package qlogwriter + +import ( + "bytes" + "encoding/json" + "io" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +type nopWriteCloserImpl struct{ io.Writer } + +func (nopWriteCloserImpl) Close() error { return nil } + +func nopWriteCloser(w io.Writer) io.WriteCloser { + return &nopWriteCloserImpl{Writer: w} +} + +func unmarshal(data []byte, v any) error { + if bytes.Equal(data[:1], recordSeparator) { + data = data[1:] + } + return json.Unmarshal(data, v) +} + +func TestTraceMetadata(t *testing.T) { + t.Run("non-connection trace", func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + buf := &bytes.Buffer{} + trace := NewFileSeq(nopWriteCloser(buf)) + go trace.Run() + producer := trace.AddProducer() + producer.Close() + + testTraceMetadata(t, buf, "transport", "", []string{}) + }) + }) + + t.Run("connection trace", func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + buf := &bytes.Buffer{} + trace := NewConnectionFileSeq( + nopWriteCloser(buf), + false, + protocol.ParseConnectionID([]byte{0xde, 0xad, 0xbe, 0xef}), + []string{"urn:ietf:params:qlog:events:foo", "urn:ietf:params:qlog:events:bar"}, + ) + + require.False(t, trace.SupportsSchemas("urn:ietf:params:qlog:events:baz")) + require.True(t, trace.SupportsSchemas("urn:ietf:params:qlog:events:foo")) + require.True(t, trace.SupportsSchemas("urn:ietf:params:qlog:events:bar")) + + go trace.Run() + producer := trace.AddProducer() + producer.Close() + + testTraceMetadata(t, + buf, + "server", + "deadbeef", + []string{"urn:ietf:params:qlog:events:foo", "urn:ietf:params:qlog:events:bar"}, + ) + }) + }) +} + +func testTraceMetadata(t *testing.T, + buf *bytes.Buffer, + expectedVantagePoint, + expectedGroupID string, + expectedEventSchemas []string, +) { + var m map[string]any + require.NoError(t, unmarshal(buf.Bytes(), &m)) + require.Equal(t, "0.3", m["qlog_version"]) + require.Contains(t, m, "title") + require.Contains(t, m, "trace") + tr := m["trace"].(map[string]any) + require.Contains(t, tr, "common_fields") + commonFields := tr["common_fields"].(map[string]any) + if expectedGroupID != "" { + require.Contains(t, commonFields, "group_id") + require.Equal(t, expectedGroupID, commonFields["group_id"]) + } else { + require.NotContains(t, commonFields, "group_id") + } + require.Contains(t, commonFields, "reference_time") + referenceTimeMap := commonFields["reference_time"].(map[string]any) + require.Contains(t, referenceTimeMap, "clock_type") + require.Equal(t, "monotonic", referenceTimeMap["clock_type"]) + require.Contains(t, referenceTimeMap, "epoch") + require.Equal(t, "unknown", referenceTimeMap["epoch"]) + require.Contains(t, referenceTimeMap, "wall_clock_time") + wallClockTimeStr := referenceTimeMap["wall_clock_time"].(string) + wallClockTime, err := time.Parse(time.RFC3339Nano, wallClockTimeStr) + require.NoError(t, err) + require.Equal(t, time.Now().UTC(), wallClockTime.UTC()) + require.Contains(t, tr, "vantage_point") + vantagePoint := tr["vantage_point"].(map[string]any) + require.Equal(t, expectedVantagePoint, vantagePoint["type"]) + if len(expectedEventSchemas) > 0 { + require.Contains(t, tr, "event_schemas") + eventSchemas := tr["event_schemas"].([]any) + for i, schema := range eventSchemas { + require.Equal(t, expectedEventSchemas[i], schema) + } + } else { + require.NotContains(t, tr, "event_schemas") + } +} diff --git a/third_party/quic-go/qlogwriter/writer.go b/third_party/quic-go/qlogwriter/writer.go new file mode 100644 index 0000000..9efc18c --- /dev/null +++ b/third_party/quic-go/qlogwriter/writer.go @@ -0,0 +1,229 @@ +package qlogwriter + +import ( + "bytes" + "fmt" + "io" + "log" + "slices" + "sync" + "time" + + "github.com/apernet/quic-go/qlogwriter/jsontext" +) + +// Trace represents a qlog trace that can have multiple event producers. +// Each producer can record events to the trace independently. +// When the last producer is closed, the underlying trace is closed as well. +type Trace interface { + // AddProducer creates a new Recorder for this trace. + // Each Recorder can record events independently. + AddProducer() Recorder + + // SupportsSchemas returns true if the trace supports the given schema. + SupportsSchemas(schema string) bool +} + +// Recorder is used to record events to a qlog trace. +// It is safe for concurrent use by multiple goroutines. +type Recorder interface { + // RecordEvent records a single Event to the trace. + // It must not be called after Close. + RecordEvent(Event) + // Close signals that this producer is done recording events. + // When all producers are closed, the underlying trace is closed. + // It must not be called concurrently with RecordEvent. + io.Closer +} + +// Event represents a qlog event that can be encoded to JSON. +// Each event must provide its name and a method to encode itself using a jsontext.Encoder. +type Event interface { + // Name returns the name of the event, as it should appear in the qlog output + Name() string + // Encode writes the event's data to the provided jsontext.Encoder + Encode(encoder *jsontext.Encoder, eventTime time.Time) error +} + +// RecordSeparator is the record separator byte for the JSON-SEQ format +const RecordSeparator byte = 0x1e + +var recordSeparator = []byte{RecordSeparator} + +type event struct { + Time time.Time + Event Event +} + +const eventChanSize = 50 + +// FileSeq represents a qlog trace using the JSON-SEQ format, +// https://www.ietf.org/archive/id/draft-ietf-quic-qlog-main-schema-12.html#section-5 +// qlog event producers can be created by calling AddProducer. +// The underlying io.WriteCloser is closed when the last producer is removed. +type FileSeq struct { + w io.WriteCloser + enc *jsontext.Encoder + referenceTime time.Time + + runStopped chan struct{} + encodeErr error + events chan event + done chan struct{} + + mx sync.Mutex + producers int + closed bool + + eventSchemas []string +} + +var _ Trace = &FileSeq{} + +// NewFileSeq creates a new JSON-SEQ qlog trace to log transport events. +func NewFileSeq(w io.WriteCloser) *FileSeq { + return newFileSeq(w, "transport", nil, nil) +} + +// NewConnectionFileSeq creates a new qlog trace to log connection events. +func NewConnectionFileSeq(w io.WriteCloser, isClient bool, odcid ConnectionID, eventSchemas []string) *FileSeq { + pers := "server" + if isClient { + pers = "client" + } + return newFileSeq(w, pers, &odcid, eventSchemas) +} + +func newFileSeq(w io.WriteCloser, pers string, odcid *ConnectionID, eventSchemas []string) *FileSeq { + now := time.Now() + buf := &bytes.Buffer{} + enc := jsontext.NewEncoder(buf) + if _, err := buf.Write(recordSeparator); err != nil { + panic(fmt.Sprintf("qlog encoding into a bytes.Buffer failed: %s", err)) + } + if err := (&traceHeader{ + VantagePointType: pers, + GroupID: odcid, + ReferenceTime: now, + EventSchemas: eventSchemas, + }).Encode(enc); err != nil { + panic(fmt.Sprintf("qlog encoding into a bytes.Buffer failed: %s", err)) + } + _, encodeErr := w.Write(buf.Bytes()) + + return &FileSeq{ + w: w, + referenceTime: now, + enc: jsontext.NewEncoder(w), + runStopped: make(chan struct{}), + encodeErr: encodeErr, + events: make(chan event, eventChanSize), + done: make(chan struct{}), + eventSchemas: eventSchemas, + } +} + +func (t *FileSeq) SupportsSchemas(schema string) bool { + return slices.Contains(t.eventSchemas, schema) +} + +func (t *FileSeq) AddProducer() Recorder { + t.mx.Lock() + defer t.mx.Unlock() + if t.closed { + return nil + } + + t.producers++ + + return &Writer{t: t} +} + +func (t *FileSeq) record(eventTime time.Time, details Event) { + t.mx.Lock() + + if t.closed { + t.mx.Unlock() + return + } + t.mx.Unlock() + + t.events <- event{Time: eventTime, Event: details} +} + +func (t *FileSeq) Run() { + defer close(t.runStopped) + + for { + select { + case <-t.done: + for { + select { + case e := <-t.events: + t.encodeEvent(e) + default: + if t.encodeErr != nil { + log.Printf("exporting qlog failed: %s\n", t.encodeErr) + } + return + } + } + case e := <-t.events: + t.encodeEvent(e) + } + } +} + +func (t *FileSeq) encodeEvent(e event) { + if t.encodeErr != nil { + return + } + if _, err := t.w.Write(recordSeparator); err != nil { + t.encodeErr = err + return + } + h := encoderHelper{enc: t.enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("time")) + h.WriteToken(jsontext.Float(float64(e.Time.Sub(t.referenceTime).Nanoseconds()) / 1e6)) + h.WriteToken(jsontext.String("name")) + h.WriteToken(jsontext.String(e.Event.Name())) + h.WriteToken(jsontext.String("data")) + if err := e.Event.Encode(t.enc, e.Time); err != nil { + t.encodeErr = err + return + } + h.WriteToken(jsontext.EndObject) + if h.err != nil { + t.encodeErr = h.err + } +} + +func (t *FileSeq) removeProducer() { + t.mx.Lock() + t.producers-- + last := t.producers == 0 + if last { + t.closed = true + } + t.mx.Unlock() + + if last { + close(t.done) + <-t.runStopped // wait for Run to drain and exit + _ = t.w.Close() + } +} + +type Writer struct { + t *FileSeq +} + +func (w *Writer) Close() error { + w.t.removeProducer() + return nil +} + +func (w *Writer) RecordEvent(ev Event) { + w.t.record(time.Now(), ev) +} diff --git a/third_party/quic-go/qlogwriter/writer_test.go b/third_party/quic-go/qlogwriter/writer_test.go new file mode 100644 index 0000000..3e6fe65 --- /dev/null +++ b/third_party/quic-go/qlogwriter/writer_test.go @@ -0,0 +1,116 @@ +package qlogwriter + +import ( + "bytes" + "errors" + "fmt" + "io" + "log" + "os" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/qlogwriter/jsontext" + + "github.com/stretchr/testify/require" +) + +type testEvent struct { + message string +} + +func (e testEvent) Name() string { + return "transport:test_event" +} + +func (e testEvent) Encode(enc *jsontext.Encoder, _ time.Time) error { + h := encoderHelper{enc: enc} + h.WriteToken(jsontext.BeginObject) + h.WriteToken(jsontext.String("message")) + h.WriteToken(jsontext.String(e.message)) + h.WriteToken(jsontext.EndObject) + return h.err +} + +type limitedWriter struct { + io.WriteCloser + N int + written int +} + +func (w *limitedWriter) Write(p []byte) (int, error) { + if w.written+len(p) > w.N { + return 0, errors.New("writer full") + } + n, err := w.WriteCloser.Write(p) + w.written += n + return n, err +} + +func TestWritingStopping(t *testing.T) { + buf := &bytes.Buffer{} + fileSeq := NewFileSeq(&limitedWriter{WriteCloser: nopWriteCloser(buf), N: 250}) + writer := fileSeq.AddProducer() + go fileSeq.Run() + + for i := range 1000 { + writer.RecordEvent(testEvent{message: fmt.Sprintf("test message %d", i)}) + } + + var logBuf bytes.Buffer + log.SetOutput(&logBuf) + defer log.SetOutput(os.Stdout) + + writer.Close() + + require.Contains(t, logBuf.String(), "writer full") + + // events after closing are ignored + logBuf.Reset() + writer.RecordEvent(testEvent{message: "foobar"}) + require.Empty(t, logBuf.String()) +} + +type blockingWriter struct { + bytes.Buffer + block bool + unblock chan struct{} +} + +func (w *blockingWriter) Write(b []byte) (int, error) { + if w.block { + <-w.unblock + } + return w.Buffer.Write(b) +} + +// TestRecordCloseRace triggers a race between record and Close. +func TestRecordCloseRace(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + w := &blockingWriter{unblock: make(chan struct{})} + trace := NewFileSeq(nopWriteCloser(w)) + go trace.Run() + synctest.Wait() // Run is blocked waiting for events + + producer := trace.AddProducer() + require.NotNil(t, producer) + + w.block = true + const numEvents = eventChanSize + 1 + for i := range numEvents { + producer.RecordEvent(testEvent{message: fmt.Sprintf("event %d", i)}) + } + + go producer.RecordEvent(testEvent{message: "last event"}) + synctest.Wait() // goroutine is blocked on full channel + + close(w.unblock) // let Run() finish + producer.Close() + + for i := range numEvents { + require.Contains(t, w.String(), fmt.Sprintf(`"message":"event %d"`, i)) + } + require.Contains(t, w.String(), `"message":"last event"`) + }) +} diff --git a/third_party/quic-go/quic_linux_test.go b/third_party/quic-go/quic_linux_test.go new file mode 100644 index 0000000..ec5c69f --- /dev/null +++ b/third_party/quic-go/quic_linux_test.go @@ -0,0 +1,12 @@ +//go:build linux + +package quic + +import ( + "fmt" +) + +func init() { + major, minor := kernelVersion() + fmt.Printf("Kernel Version: %d.%d\n\n", major, minor) +} diff --git a/third_party/quic-go/quic_test.go b/third_party/quic-go/quic_test.go new file mode 100644 index 0000000..9f6eaf1 --- /dev/null +++ b/third_party/quic-go/quic_test.go @@ -0,0 +1,87 @@ +package quic + +import ( + "bytes" + "fmt" + "net" + "os" + "runtime/pprof" + "strconv" + "strings" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +// in the tests for the stream deadlines we set a deadline +// and wait to make an assertion when Read / Write was unblocked +// on the CIs, the timing is a lot less precise, so scale every duration by this factor +func scaleDuration(t time.Duration) time.Duration { + scaleFactor := 1 + if f, err := strconv.Atoi(os.Getenv("TIMESCALE_FACTOR")); err == nil { // parsing "" errors, so this works fine if the env is not set + scaleFactor = f + } + if scaleFactor == 0 { + panic("TIMESCALE_FACTOR is 0") + } + return time.Duration(scaleFactor) * t +} + +func newUDPConnLocalhost(t testing.TB) *net.UDPConn { + t.Helper() + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + t.Cleanup(func() { conn.Close() }) + return conn +} + +func getPacket(t *testing.T, connID protocol.ConnectionID) []byte { + return getPacketWithPacketType(t, connID, protocol.PacketTypeHandshake, 2) +} + +func getPacketWithPacketType(t *testing.T, connID protocol.ConnectionID, typ protocol.PacketType, length protocol.ByteCount) []byte { + t.Helper() + b, err := (&wire.ExtendedHeader{ + Header: wire.Header{ + Type: typ, + DestConnectionID: connID, + Length: length, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen2, + }).Append(nil, protocol.Version1) + require.NoError(t, err) + return append(b, bytes.Repeat([]byte{42}, int(length)-2)...) +} + +func areConnsRunning() bool { + var b bytes.Buffer + pprof.Lookup("goroutine").WriteTo(&b, 1) + return strings.Contains(b.String(), "quic-go.(*connection).run") +} + +func areTransportsRunning() bool { + var b bytes.Buffer + pprof.Lookup("goroutine").WriteTo(&b, 1) + return strings.Contains(b.String(), "quic-go.(*Transport).listen") +} + +func TestMain(m *testing.M) { + status := m.Run() + if status != 0 { + os.Exit(status) + } + if areConnsRunning() { + fmt.Println("stray connection goroutines found") + os.Exit(1) + } + if areTransportsRunning() { + fmt.Println("stray transport goroutines found") + os.Exit(1) + } + os.Exit(status) +} diff --git a/third_party/quic-go/quicvarint/io.go b/third_party/quic-go/quicvarint/io.go new file mode 100644 index 0000000..8ea10ac --- /dev/null +++ b/third_party/quic-go/quicvarint/io.go @@ -0,0 +1,98 @@ +package quicvarint + +import ( + "bytes" + "io" +) + +// Reader implements both the io.ByteReader and io.Reader interfaces. +type Reader interface { + io.ByteReader + io.Reader +} + +var _ Reader = &bytes.Reader{} + +// A Peeker can peek bytes without consuming them. +type Peeker interface { + Peek(b []byte) (int, error) +} + +// Peek reads a number in the QUIC varint format without consuming bytes. +func Peek(p Peeker) (uint64, error) { + var b [8]byte + + // first peek 1 byte to determine the varint length + if _, err := p.Peek(b[:1]); err != nil { + return 0, err + } + + l := 1 << (b[0] >> 6) // 1, 2, 4, or 8 bytes + if l == 1 { + return uint64(b[0] & 0b00111111), nil + } + if _, err := p.Peek(b[:l]); err != nil { + return 0, err + } + val, _, err := Parse(b[:l]) + return val, err +} + +type byteReader struct { + io.Reader +} + +var _ Reader = &byteReader{} + +// NewReader returns a Reader for r. +// If r already implements both io.ByteReader and io.Reader, NewReader returns r. +// Otherwise, r is wrapped to add the missing interfaces. +func NewReader(r io.Reader) Reader { + if r, ok := r.(Reader); ok { + return r + } + return &byteReader{r} +} + +func (r *byteReader) ReadByte() (byte, error) { + var b [1]byte + var n int + var err error + for n == 0 && err == nil { + n, err = r.Read(b[:]) + } + + if n == 1 && err == io.EOF { + err = nil + } + return b[0], err +} + +// Writer implements both the io.ByteWriter and io.Writer interfaces. +type Writer interface { + io.ByteWriter + io.Writer +} + +var _ Writer = &bytes.Buffer{} + +type byteWriter struct { + io.Writer +} + +var _ Writer = &byteWriter{} + +// NewWriter returns a Writer for w. +// If w already implements both io.ByteWriter and io.Writer, NewWriter returns w. +// Otherwise, w is wrapped to add the missing interfaces. +func NewWriter(w io.Writer) Writer { + if w, ok := w.(Writer); ok { + return w + } + return &byteWriter{w} +} + +func (w *byteWriter) WriteByte(c byte) error { + _, err := w.Write([]byte{c}) + return err +} diff --git a/third_party/quic-go/quicvarint/io_test.go b/third_party/quic-go/quicvarint/io_test.go new file mode 100644 index 0000000..9b58773 --- /dev/null +++ b/third_party/quic-go/quicvarint/io_test.go @@ -0,0 +1,162 @@ +package quicvarint + +import ( + "bytes" + "fmt" + "io" + "testing" + + "github.com/stretchr/testify/require" +) + +type nopReader struct{} + +func (r *nopReader) Read(_ []byte) (int, error) { + return 0, io.ErrUnexpectedEOF +} + +var _ io.Reader = &nopReader{} + +type nopWriter struct{} + +func (r *nopWriter) Write(_ []byte) (int, error) { + return 0, io.ErrShortBuffer +} + +// eofReader is a reader that returns data and the io.EOF at the same time in the last Read call +type eofReader struct { + Data []byte + pos int +} + +func (r *eofReader) Read(b []byte) (int, error) { + n := copy(b, r.Data[r.pos:]) + r.pos += n + if r.pos >= len(r.Data) { + return n, io.EOF + } + return n, nil +} + +var _ io.Writer = &nopWriter{} + +func TestReaderPassesThroughUnchanged(t *testing.T) { + b := bytes.NewReader([]byte{0}) + r := NewReader(b) + require.Equal(t, b, r) +} + +func TestReaderWrapsIOReader(t *testing.T) { + n := &nopReader{} + r := NewReader(n) + require.NotEqual(t, n, r) +} + +func TestReaderFailure(t *testing.T) { + r := NewReader(&nopReader{}) + val, err := r.ReadByte() + require.Equal(t, io.ErrUnexpectedEOF, err) + require.Equal(t, byte(0), val) +} + +func TestReaderHandlesEOF(t *testing.T) { + // test that the eofReader behaves as we expect + r := &eofReader{Data: []byte("foobar")} + b := make([]byte, 3) + n, err := r.Read(b) + require.Equal(t, 3, n) + require.NoError(t, err) + require.Equal(t, "foo", string(b)) + n, err = r.Read(b) + require.Equal(t, 3, n) + require.Equal(t, io.EOF, err) + require.Equal(t, "bar", string(b)) + n, err = r.Read(b) + require.Equal(t, io.EOF, err) + require.Zero(t, n) + + // now test using it to read varints + reader := NewReader(&eofReader{Data: Append(nil, 1337)}) + n2, err := Read(reader) + require.NoError(t, err) + require.EqualValues(t, 1337, n2) +} + +// Regression test: empty reads were being converted to successful +// reads of a zero value. +func TestReaderHandlesEmptyRead(t *testing.T) { + r, w := io.Pipe() + + go func() { + // io.Pipe turns empty writes into empty reads. + w.Write(nil) + w.Close() + }() + + br := NewReader(r) + _, err := Read(br) + require.ErrorIs(t, err, io.EOF) +} + +func TestWriterPassesThroughUnchanged(t *testing.T) { + b := &bytes.Buffer{} + w := NewWriter(b) + require.Equal(t, b, w) +} + +func TestWriterWrapsIOWriter(t *testing.T) { + n := &nopWriter{} + w := NewWriter(n) + require.NotEqual(t, n, w) +} + +func TestWriterFailure(t *testing.T) { + w := NewWriter(&nopWriter{}) + err := w.WriteByte(0) + require.Equal(t, io.ErrShortBuffer, err) +} + +type bufPeeker []byte + +func (p bufPeeker) Peek(b []byte) (int, error) { + if len(p) < len(b) { + return copy(b, p), io.ErrUnexpectedEOF + } + return copy(b, p), nil +} + +func TestPeek(t *testing.T) { + for _, c := range []bufPeeker{ + {0b00011001}, // 1-byte + {0b01111011, 0xbd}, // 2-byte + {0b10011101, 0x7f, 0x3e, 0x7d}, // 4-byte + {0xc2, 0x19, 0x7c, 0x5e, 0xff, 0x14, 0xe8, 0x8c}, // 8-byte + } { + t.Run(fmt.Sprintf("%d bytes", len(c)), func(t *testing.T) { + peekVal, err := Peek(append(c, []byte("foobar")...)) // append some data, which doesn't matter + require.NoError(t, err) + parseVal, _, err := Parse(c) + require.NoError(t, err) + require.Equal(t, parseVal, peekVal) + }) + } +} + +func TestPeekErrors(t *testing.T) { + errorCases := []struct { + name string + input bufPeeker + }{ + {"empty input", bufPeeker{}}, + {"2-byte, missing 1", bufPeeker{0b01000001}}, + {"4-byte, missing 1", bufPeeker{0b10000000, 0, 0}}, + {"8-byte, missing 1", bufPeeker{0b11000000, 0, 0, 0, 0, 0, 0}}, + } + + for _, tc := range errorCases { + t.Run(tc.name, func(t *testing.T) { + _, err := Peek(tc.input) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + }) + } +} diff --git a/third_party/quic-go/quicvarint/varint.go b/third_party/quic-go/quicvarint/varint.go new file mode 100644 index 0000000..52fb153 --- /dev/null +++ b/third_party/quic-go/quicvarint/varint.go @@ -0,0 +1,180 @@ +package quicvarint + +import ( + "encoding/binary" + "fmt" + "io" +) + +// taken from the QUIC draft +const ( + // Min is the minimum value allowed for a QUIC varint. + Min = 0 + + // Max is the maximum allowed value for a QUIC varint (2^62-1). + Max = maxVarInt8 + + maxVarInt1 = 63 + maxVarInt2 = 16383 + maxVarInt4 = 1073741823 + maxVarInt8 = 4611686018427387903 +) + +type varintLengthError struct { + Num uint64 +} + +func (e *varintLengthError) Error() string { + return fmt.Sprintf("value doesn't fit into 62 bits: %d", e.Num) +} + +// Read reads a number in the QUIC varint format from r. +func Read(r io.ByteReader) (uint64, error) { + firstByte, err := r.ReadByte() + if err != nil { + return 0, err + } + // the first two bits of the first byte encode the length + l := 1 << ((firstByte & 0xc0) >> 6) + b1 := firstByte & (0xff - 0xc0) + if l == 1 { + return uint64(b1), nil + } + b2, err := r.ReadByte() + if err != nil { + return 0, err + } + if l == 2 { + return uint64(b2) + uint64(b1)<<8, nil + } + b3, err := r.ReadByte() + if err != nil { + return 0, err + } + b4, err := r.ReadByte() + if err != nil { + return 0, err + } + if l == 4 { + return uint64(b4) + uint64(b3)<<8 + uint64(b2)<<16 + uint64(b1)<<24, nil + } + b5, err := r.ReadByte() + if err != nil { + return 0, err + } + b6, err := r.ReadByte() + if err != nil { + return 0, err + } + b7, err := r.ReadByte() + if err != nil { + return 0, err + } + b8, err := r.ReadByte() + if err != nil { + return 0, err + } + return uint64(b8) + uint64(b7)<<8 + uint64(b6)<<16 + uint64(b5)<<24 + uint64(b4)<<32 + uint64(b3)<<40 + uint64(b2)<<48 + uint64(b1)<<56, nil +} + +// Parse reads a number in the QUIC varint format. +// It returns the number of bytes consumed. +func Parse(b []byte) (uint64 /* value */, int /* bytes consumed */, error) { + if len(b) == 0 { + return 0, 0, io.EOF + } + + first := b[0] + switch first >> 6 { + case 0: // 1-byte encoding: 00xxxxxx + return uint64(first & 0b00111111), 1, nil + case 1: // 2-byte encoding: 01xxxxxx + if len(b) < 2 { + return 0, 0, io.ErrUnexpectedEOF + } + return uint64(b[1]) | uint64(first&0b00111111)<<8, 2, nil + case 2: // 4-byte encoding: 10xxxxxx + if len(b) < 4 { + return 0, 0, io.ErrUnexpectedEOF + } + return uint64(b[3]) | uint64(b[2])<<8 | uint64(b[1])<<16 | uint64(first&0b00111111)<<24, 4, nil + case 3: // 8-byte encoding: 00xxxxxx + if len(b) < 8 { + return 0, 0, io.ErrUnexpectedEOF + } + // binary.BigEndian.Uint64 only reads the first 8 bytes. Passing the full slice avoids slicing overhead. + return binary.BigEndian.Uint64(b) & 0x3fffffffffffffff, 8, nil + } + + panic("unreachable") +} + +// Append appends i in the QUIC varint format. +func Append(b []byte, i uint64) []byte { + if i <= maxVarInt1 { + return append(b, uint8(i)) + } + if i <= maxVarInt2 { + return append(b, []byte{uint8(i>>8) | 0x40, uint8(i)}...) + } + if i <= maxVarInt4 { + return append(b, []byte{uint8(i>>24) | 0x80, uint8(i >> 16), uint8(i >> 8), uint8(i)}...) + } + if i <= maxVarInt8 { + return append(b, []byte{ + uint8(i>>56) | 0xc0, uint8(i >> 48), uint8(i >> 40), uint8(i >> 32), + uint8(i >> 24), uint8(i >> 16), uint8(i >> 8), uint8(i), + }...) + } + panic(&varintLengthError{Num: i}) +} + +// AppendWithLen append i in the QUIC varint format with the desired length. +func AppendWithLen(b []byte, i uint64, length int) []byte { + if length != 1 && length != 2 && length != 4 && length != 8 { + panic("invalid varint length") + } + l := Len(i) + if l == length { + return Append(b, i) + } + if l > length { + panic(fmt.Sprintf("cannot encode %d in %d bytes", i, length)) + } + switch length { + case 2: + b = append(b, 0b01000000) + case 4: + b = append(b, 0b10000000) + case 8: + b = append(b, 0b11000000) + } + for range length - l - 1 { + b = append(b, 0) + } + for j := range l { + b = append(b, uint8(i>>(8*(l-1-j)))) + } + return b +} + +// Len determines the number of bytes that will be needed to write the number i. +// +//gcassert:inline +func Len(i uint64) int { + if i <= maxVarInt1 { + return 1 + } + if i <= maxVarInt2 { + return 2 + } + if i <= maxVarInt4 { + return 4 + } + if i <= maxVarInt8 { + return 8 + } + // Don't use a fmt.Sprintf here to format the error message. + // The function would then exceed the inlining budget. + panic(&varintLengthError{Num: i}) +} diff --git a/third_party/quic-go/quicvarint/varint_test.go b/third_party/quic-go/quicvarint/varint_test.go new file mode 100644 index 0000000..7c0116e --- /dev/null +++ b/third_party/quic-go/quicvarint/varint_test.go @@ -0,0 +1,350 @@ +package quicvarint + +import ( + "bytes" + "fmt" + "io" + "math/rand/v2" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestLimits(t *testing.T) { + require.Equal(t, 0, Min) + require.Equal(t, uint64(1<<62-1), uint64(Max)) +} + +func TestRead(t *testing.T) { + tests := []struct { + name string + input []byte + expected uint64 + }{ + {"1 byte", []byte{0b00011001}, 25}, + {"2 byte", []byte{0b01111011, 0xbd}, 15293}, + {"4 byte", []byte{0b10011101, 0x7f, 0x3e, 0x7d}, 494878333}, + {"8 byte", []byte{0b11000010, 0x19, 0x7c, 0x5e, 0xff, 0x14, 0xe8, 0x8c}, 151288809941952652}, + {"too long", []byte{0b01000000, 0x25}, 37}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + b := bytes.NewReader(tt.input) + val, err := Read(b) + require.NoError(t, err) + require.Equal(t, tt.expected, val) + require.Zero(t, b.Len()) + }) + } +} + +func TestParse(t *testing.T) { + tests := []struct { + name string + input []byte + expectedValue uint64 + expectedLen int + }{ + {"1 byte", []byte{0b00011001}, 25, 1}, + {"2 byte", []byte{0b01111011, 0xbd}, 15293, 2}, + {"4 byte", []byte{0b10011101, 0x7f, 0x3e, 0x7d}, 494878333, 4}, + {"8 byte", []byte{0b11000010, 0x19, 0x7c, 0x5e, 0xff, 0x14, 0xe8, 0x8c}, 151288809941952652, 8}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + value, l, err := Parse(tt.input) + require.Equal(t, tt.expectedValue, value) + require.Equal(t, tt.expectedLen, l) + require.Nil(t, err) + }) + } +} + +func TestParsingFailures(t *testing.T) { + tests := []struct { + name string + input []byte + expectedErr error + }{ + { + name: "empty slice", + input: []byte{}, + expectedErr: io.EOF, + }, + { + name: "2-byte encoding: not enough bytes", + input: []byte{0b01000001}, + expectedErr: io.ErrUnexpectedEOF, + }, + { + name: "4-byte encoding: not enough bytes", + input: []byte{0b10000000, 0x0, 0x0}, + expectedErr: io.ErrUnexpectedEOF, + }, + { + name: "8-byte encoding: not enough bytes", + input: []byte{0b11000000, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0}, + expectedErr: io.ErrUnexpectedEOF, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + value, l, err := Parse(tt.input) + require.Equal(t, uint64(0), value) + require.Equal(t, 0, l) + require.Equal(t, tt.expectedErr, err) + }) + } +} + +func TestVarintEncoding(t *testing.T) { + tests := []struct { + name string + value uint64 + expected []byte + }{ + {"1 byte number", 37, []byte{0x25}}, + {"maximum 1 byte number", maxVarInt1, []byte{0b00111111}}, + {"minimum 2 byte number", maxVarInt1 + 1, []byte{0x40, maxVarInt1 + 1}}, + {"2 byte number", 15293, []byte{0b01000000 ^ 0x3b, 0xbd}}, + {"maximum 2 byte number", maxVarInt2, []byte{0b01111111, 0xff}}, + {"minimum 4 byte number", maxVarInt2 + 1, []byte{0b10000000, 0, 0x40, 0}}, + {"4 byte number", 494878333, []byte{0b10000000 ^ 0x1d, 0x7f, 0x3e, 0x7d}}, + {"maximum 4 byte number", maxVarInt4, []byte{0b10111111, 0xff, 0xff, 0xff}}, + {"minimum 8 byte number", maxVarInt4 + 1, []byte{0b11000000, 0, 0, 0, 0x40, 0, 0, 0}}, + {"8 byte number", 151288809941952652, []byte{0xc2, 0x19, 0x7c, 0x5e, 0xff, 0x14, 0xe8, 0x8c}}, + {"maximum 8 byte number", maxVarInt8, []byte{0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.expected, Append(nil, tt.value)) + }) + } + + t.Run("panics when given a too large number (> 62 bit)", func(t *testing.T) { + require.PanicsWithError(t, + fmt.Sprintf("value doesn't fit into 62 bits: %d", maxVarInt8+1), + func() { Append(nil, maxVarInt8+1) }, + ) + }) +} + +func TestAppendWithLen(t *testing.T) { + tests := []struct { + name string + value uint64 + length int + expected []byte + }{ + {"1-byte number in minimal encoding", 37, 1, []byte{0x25}}, + {"1-byte number in 2 bytes", 37, 2, []byte{0b01000000, 0x25}}, + {"1-byte number in 4 bytes", 37, 4, []byte{0b10000000, 0, 0, 0x25}}, + {"1-byte number in 8 bytes", 37, 8, []byte{0b11000000, 0, 0, 0, 0, 0, 0, 0x25}}, + {"2-byte number in 4 bytes", 15293, 4, []byte{0b10000000, 0, 0x3b, 0xbd}}, + {"4-byte number in 8 bytes", 494878333, 8, []byte{0b11000000, 0, 0, 0, 0x1d, 0x7f, 0x3e, 0x7d}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + b := AppendWithLen(nil, tt.value, tt.length) + require.Equal(t, tt.expected, b) + + if tt.length > 1 { + v, n, err := Parse(b) + require.NoError(t, err) + require.Equal(t, tt.length, n) + require.Equal(t, tt.value, v) + } + }) + } +} + +func TestAppendWithLenFailures(t *testing.T) { + tests := []struct { + name string + value uint64 + length int + }{ + {"invalid length", 25, 3}, + {"too short for 2 bytes", maxVarInt1 + 1, 1}, + {"too short for 4 bytes", maxVarInt2 + 1, 2}, + {"too short for 8 bytes", maxVarInt4 + 1, 4}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Panics(t, func() { + AppendWithLen(nil, tt.value, tt.length) + }) + }) + } +} + +func TestLen(t *testing.T) { + tests := []struct { + name string + input uint64 + expected int + }{ + {"zero", 0, 1}, + {"max 1 byte", maxVarInt1, 1}, + {"min 2 bytes", maxVarInt1 + 1, 2}, + {"max 2 bytes", maxVarInt2, 2}, + {"min 4 bytes", maxVarInt2 + 1, 4}, + {"max 4 bytes", maxVarInt4, 4}, + {"min 8 bytes", maxVarInt4 + 1, 8}, + {"max 8 bytes", maxVarInt8, 8}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.expected, Len(tt.input)) + }) + } + + t.Run("panics on too large number", func(t *testing.T) { + require.PanicsWithError(t, + fmt.Sprintf("value doesn't fit into 62 bits: %d", maxVarInt8+1), + func() { Len(maxVarInt8 + 1) }, + ) + }) +} + +type benchmarkValue struct { + b []byte + v uint64 +} + +func randomValues(maxValue uint64) []benchmarkValue { + r := rand.New(rand.NewPCG(13, 37)) + + const num = 1025 + bv := make([]benchmarkValue, num) + for i := range num { + v := r.Uint64() % maxValue + bv[i].v = v + bv[i].b = Append([]byte{}, v) + } + return bv +} + +// using a reader that is also an io.ByteReader +func BenchmarkReadBytesReader(b *testing.B) { + b.Run("1-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt1), false) }) + b.Run("2-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt2), false) }) + b.Run("4-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt4), false) }) + b.Run("8-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt8), false) }) +} + +// using a reader that is not an io.ByteReader +func BenchmarkReadSimpleReader(b *testing.B) { + b.Run("1-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt1), true) }) + b.Run("2-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt2), true) }) + b.Run("4-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt4), true) }) + b.Run("8-byte", func(b *testing.B) { benchmarkRead(b, randomValues(maxVarInt8), true) }) +} + +// simpleReader satisfies io.Reader, but not io.ByteReader +// This means that NewReader will need to wrap the reader. +type simpleReader struct { + io.Reader +} + +func benchmarkRead(b *testing.B, inputs []benchmarkValue, wrapBytesReader bool) { + r := bytes.NewReader([]byte{}) + var vr Reader + if wrapBytesReader { + vr = NewReader(&simpleReader{r}) + } else { + vr = NewReader(r) + } + + var i int + for b.Loop() { + index := i % len(inputs) + i++ + r.Reset(inputs[index].b) + val, err := Read(vr) + if err != nil { + b.Fatal(err) + } + if val != inputs[index].v { + b.Fatalf("expected %d, got %d", inputs[index].v, val) + } + } +} + +func BenchmarkParse(b *testing.B) { + b.Run("1-byte", func(b *testing.B) { benchmarkParse(b, randomValues(maxVarInt1)) }) + b.Run("2-byte", func(b *testing.B) { benchmarkParse(b, randomValues(maxVarInt2)) }) + b.Run("4-byte", func(b *testing.B) { benchmarkParse(b, randomValues(maxVarInt4)) }) + b.Run("8-byte", func(b *testing.B) { benchmarkParse(b, randomValues(maxVarInt8)) }) +} + +func benchmarkParse(b *testing.B, inputs []benchmarkValue) { + var i int + for b.Loop() { + index := i % len(inputs) + i++ + val, n, err := Parse(inputs[index].b) + if err != nil { + b.Fatal(err) + } + if n != len(inputs[index].b) { + b.Fatalf("expected to consume %d bytes, consumed %d", len(inputs[i].b), n) + } + if val != inputs[index].v { + b.Fatalf("expected %d, got %d", inputs[index].v, val) + } + } +} + +func BenchmarkAppend(b *testing.B) { + b.Run("1-byte", func(b *testing.B) { benchmarkAppend(b, randomValues(maxVarInt1)) }) + b.Run("2-byte", func(b *testing.B) { benchmarkAppend(b, randomValues(maxVarInt2)) }) + b.Run("4-byte", func(b *testing.B) { benchmarkAppend(b, randomValues(maxVarInt4)) }) + b.Run("8-byte", func(b *testing.B) { benchmarkAppend(b, randomValues(maxVarInt8)) }) +} + +func benchmarkAppend(b *testing.B, inputs []benchmarkValue) { + buf := make([]byte, 8) + + var i int + for b.Loop() { + buf = buf[:0] + index := i % len(inputs) + i++ + buf = Append(buf, inputs[index].v) + + if !bytes.Equal(buf, inputs[index].b) { + b.Fatalf("expected to write %v, wrote %v", inputs[index].b, buf) + } + } +} + +func BenchmarkAppendWithLen(b *testing.B) { + b.Run("1-byte", func(b *testing.B) { benchmarkAppendWithLen(b, randomValues(maxVarInt1)) }) + b.Run("2-byte", func(b *testing.B) { benchmarkAppendWithLen(b, randomValues(maxVarInt2)) }) + b.Run("4-byte", func(b *testing.B) { benchmarkAppendWithLen(b, randomValues(maxVarInt4)) }) + b.Run("8-byte", func(b *testing.B) { benchmarkAppendWithLen(b, randomValues(maxVarInt8)) }) +} + +func benchmarkAppendWithLen(b *testing.B, inputs []benchmarkValue) { + buf := make([]byte, 8) + + var i int + for b.Loop() { + buf = buf[:0] + index := i % len(inputs) + i++ + buf = AppendWithLen(buf, inputs[index].v, len(inputs[index].b)) + + if !bytes.Equal(buf, inputs[index].b) { + b.Fatalf("expected to write %v, wrote %v", inputs[index].b, buf) + } + } +} diff --git a/third_party/quic-go/receive_stream.go b/third_party/quic-go/receive_stream.go new file mode 100644 index 0000000..74afca7 --- /dev/null +++ b/third_party/quic-go/receive_stream.go @@ -0,0 +1,578 @@ +package quic + +import ( + "fmt" + "io" + "sync" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +// A ReceiveStream is a unidirectional Receive Stream. +type ReceiveStream struct { + mutex sync.Mutex + + streamID protocol.StreamID + + sender streamSender + + frameQueue *frameSorter + finalOffset protocol.ByteCount + receiveFinalSizeCallback func(int64) + + currentFrame []byte + currentFrameDone func() + readPosInFrame int + currentFrameIsLast bool // is the currentFrame the last frame on this stream + + queuedStopSending bool + queuedMaxStreamData bool + + // Set once we read the io.EOF or the cancellation error. + // Note that for local cancellations, this doesn't necessarily mean that we know the final offset yet. + errorRead bool + completed bool // set once we've called streamSender.onStreamCompleted + cancelledRemotely bool + cancelledLocally bool + cancelErr *StreamError + closeForShutdownErr error + + readPos protocol.ByteCount + reliableSize protocol.ByteCount + + readChan chan struct{} + readOnce chan struct{} // cap: 1, to protect against concurrent use of Read + deadline monotime.Time + + flowController *streamFlowController +} + +var ( + _ streamControlFrameGetter = &ReceiveStream{} + _ receiveStreamFrameHandler = &ReceiveStream{} +) + +func newReceiveStream( + streamID protocol.StreamID, + sender streamSender, + flowController *streamFlowController, +) *ReceiveStream { + return &ReceiveStream{ + streamID: streamID, + sender: sender, + flowController: flowController, + frameQueue: newFrameSorter(), + readChan: make(chan struct{}, 1), + readOnce: make(chan struct{}, 1), + finalOffset: protocol.MaxByteCount, + } +} + +// StreamID returns the stream ID. +func (s *ReceiveStream) StreamID() StreamID { + return s.streamID +} + +// SetReceiveFinalSizeCallback sets a callback that is called when the receive stream's final size is known. +// The final size is learned from a FIN or RESET_STREAM frame. +// Most applications don't need this. It is mainly useful for protocol layers +// that need exact stream final sizes, such as WebTransport flow control accounting. +// If the final size is already known, the callback is called before this method returns. +// When the final size is learned later, the callback is called from the connection's event loop and must not block. +// The callback is not called if the connection is closed before the final size is known. +// Setting a nil callback removes it if the final size is not yet known. +func (s *ReceiveStream) SetReceiveFinalSizeCallback(callback func(int64)) { + s.mutex.Lock() + size := s.finalOffset + + // final size is already known + if size != protocol.MaxByteCount { + s.mutex.Unlock() + if callback != nil { + callback(int64(size)) + } + return + } + + if s.closeForShutdownErr != nil { + s.mutex.Unlock() + return + } + s.receiveFinalSizeCallback = callback + s.mutex.Unlock() +} + +// Read reads data from the stream. +// Read can be made to time out using [ReceiveStream.SetReadDeadline]. +// If the stream was canceled, the error is a [StreamError]. +func (s *ReceiveStream) Read(p []byte) (int, error) { + // Concurrent use of Read is not permitted (and doesn't make any sense), + // but sometimes people do it anyway. + // Make sure that we only execute one call at any given time to avoid hard to debug failures. + s.readOnce <- struct{}{} + defer func() { <-s.readOnce }() + + s.mutex.Lock() + queuedStreamWindowUpdate, queuedConnWindowUpdate, n, err := s.readImpl(p) + completed := s.isNewlyCompleted() + s.mutex.Unlock() + + if completed { + s.sender.onStreamCompleted(s.streamID) + } + if queuedStreamWindowUpdate { + s.sender.onHasStreamControlFrame(s.streamID, s) + } + if queuedConnWindowUpdate { + s.sender.onHasConnectionData() + } + return n, err +} + +func (s *ReceiveStream) isNewlyCompleted() bool { + if s.completed { + return false + } + // We need to know the final offset (either via FIN or RESET_STREAM) for flow control accounting. + if s.finalOffset == protocol.MaxByteCount { + return false + } + // We're done with the stream if it was cancelled locally... + if s.cancelledLocally { + s.completed = true + return true + } + // ... or if the error (either io.EOF or the reset error) was read + if s.errorRead { + s.completed = true + return true + } + return false +} + +func (s *ReceiveStream) readImpl(p []byte) (hasStreamWindowUpdate bool, hasConnWindowUpdate bool, _ int, _ error) { + if s.currentFrameIsLast && s.currentFrame == nil { + s.errorRead = true + return false, false, 0, io.EOF + } + if s.cancelledLocally || s.isRemoteCancellationEffective() { + s.errorRead = true + return false, false, 0, s.cancelErr + } + if s.closeForShutdownErr != nil { + return false, false, 0, s.closeForShutdownErr + } + + var bytesRead int + var deadlineTimer *time.Timer + for bytesRead < len(p) { + if s.currentFrame == nil || s.readPosInFrame >= len(s.currentFrame) { + s.dequeueNextFrame() + } + if s.currentFrame == nil && bytesRead > 0 { + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, s.closeForShutdownErr + } + + for { + // Stop waiting on errors + if s.closeForShutdownErr != nil { + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, s.closeForShutdownErr + } + if s.cancelledLocally || s.isRemoteCancellationEffective() { + s.errorRead = true + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, s.cancelErr + } + + deadline := s.deadline + if !deadline.IsZero() && !monotime.Now().Before(deadline) { + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, errDeadline + } + + if s.currentFrame != nil || s.currentFrameIsLast { + break + } + + s.mutex.Unlock() + if deadline.IsZero() { + <-s.readChan + } else { + if deadlineTimer == nil { + deadlineTimer = time.NewTimer(monotime.Until(deadline)) + defer deadlineTimer.Stop() + } else { + deadlineTimer.Reset(monotime.Until(deadline)) + } + select { + case <-s.readChan: + case <-deadlineTimer.C: + } + } + s.mutex.Lock() + s.dequeueNextFrame() + } + + if bytesRead > len(p) { + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, fmt.Errorf("BUG: bytesRead (%d) > len(p) (%d) in stream.Read", bytesRead, len(p)) + } + if s.readPosInFrame > len(s.currentFrame) { + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, fmt.Errorf("BUG: readPosInFrame (%d) > frame.DataLen (%d) in stream.Read", s.readPosInFrame, len(s.currentFrame)) + } + m := copy(p[bytesRead:], s.currentFrame[s.readPosInFrame:]) + + // when a RESET_STREAM was received, the flow controller was already + // informed about the final offset for this stream + if !s.isRemoteCancellationEffective() { + hasStream, hasConn := s.flowController.AddBytesRead(protocol.ByteCount(m)) + if hasStream { + s.queuedMaxStreamData = true + hasStreamWindowUpdate = true + } + if hasConn { + hasConnWindowUpdate = true + } + } + + s.readPosInFrame += m + s.readPos += protocol.ByteCount(m) + bytesRead += m + + if s.isRemoteCancellationEffective() { + s.flowController.Abandon() + } + + if s.readPosInFrame >= len(s.currentFrame) && s.currentFrameIsLast { + s.currentFrame = nil + if s.currentFrameDone != nil { + s.currentFrameDone() + } + s.errorRead = true + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, io.EOF + } + } + if s.isRemoteCancellationEffective() { + s.errorRead = true + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, s.cancelErr + } + return hasStreamWindowUpdate, hasConnWindowUpdate, bytesRead, nil +} + +// isRemoteCancellationEffective returns whether the stream was cancelled remotely +// and all reliable data has been read. +func (s *ReceiveStream) isRemoteCancellationEffective() bool { + return s.cancelledRemotely && s.readPos >= s.reliableSize +} + +// Peek fills b with stream data, without consuming the stream data. +// It blocks until len(b) bytes are available, or an error occurs. +// It respects the stream deadline set by SetReadDeadline. +// If the stream ends before len(b) bytes are available, +// it returns the number of bytes peeked along with io.EOF. +func (s *ReceiveStream) Peek(b []byte) (int, error) { + if len(b) == 0 { + return 0, nil + } + + // prevent concurrent use with Read + s.readOnce <- struct{}{} + defer func() { <-s.readOnce }() + + return s.peekImpl(b) +} + +func (s *ReceiveStream) peekImpl(b []byte) (int, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + + var deadlineTimer *time.Timer + + for { + if s.currentFrameIsLast && s.currentFrame == nil { + return 0, io.EOF + } + if s.cancelledLocally || s.isRemoteCancellationEffective() { + return 0, s.cancelErr + } + if s.closeForShutdownErr != nil { + return 0, s.closeForShutdownErr + } + + deadline := s.deadline + if !deadline.IsZero() && !monotime.Now().Before(deadline) { + return 0, errDeadline + } + + if s.currentFrame == nil || s.readPosInFrame >= len(s.currentFrame) { + s.dequeueNextFrame() + } + + if s.currentFrame != nil && s.readPosInFrame < len(s.currentFrame) { + availableInCurrentFrame := len(s.currentFrame) - s.readPosInFrame + + if availableInCurrentFrame >= len(b) { + copy(b, s.currentFrame[s.readPosInFrame:]) + return len(b), nil + } + + offset := s.readPos + protocol.ByteCount(availableInCurrentFrame) + // First peek, then copy. + // This avoids copying data if there's not enough data in the queue. + if err := s.frameQueue.Peek(offset, b[availableInCurrentFrame:]); err == nil { + copy(b[:availableInCurrentFrame], s.currentFrame[s.readPosInFrame:]) + return len(b), nil + } + + if s.currentFrameIsLast { + copy(b[:availableInCurrentFrame], s.currentFrame[s.readPosInFrame:]) + return availableInCurrentFrame, io.EOF + } + + // If the stream was remotely cancelled and the request extends beyond the reliable size, + // return the data available with the cancel error (once it's all received). + if s.cancelledRemotely && s.readPos+protocol.ByteCount(len(b)) > s.reliableSize { + total := int(s.reliableSize - s.readPos) + needed := total - availableInCurrentFrame + // only return once all available data is contiguous + if needed <= 0 || s.frameQueue.Peek(offset, b[availableInCurrentFrame:total]) == nil { + copy(b[:availableInCurrentFrame], s.currentFrame[s.readPosInFrame:]) + return total, s.cancelErr + } + } + + // If the request extends beyond the stream's final offset, + // return the data available with EOF (once it's all received). + if s.readPos+protocol.ByteCount(len(b)) > s.finalOffset { + total := int(s.finalOffset - s.readPos) + needed := total - availableInCurrentFrame + // only return once all available data is contiguous + if needed <= 0 || s.frameQueue.Peek(offset, b[availableInCurrentFrame:total]) == nil { + copy(b[:availableInCurrentFrame], s.currentFrame[s.readPosInFrame:]) + return total, io.EOF + } + } + } + + if s.currentFrameIsLast || s.readPos >= s.finalOffset { + return 0, io.EOF + } + + s.mutex.Unlock() + if deadline.IsZero() { + <-s.readChan + } else { + if deadlineTimer == nil { + deadlineTimer = time.NewTimer(monotime.Until(deadline)) + defer deadlineTimer.Stop() + } else { + deadlineTimer.Reset(monotime.Until(deadline)) + } + select { + case <-s.readChan: + case <-deadlineTimer.C: + } + } + s.mutex.Lock() + if s.currentFrame == nil || s.readPosInFrame >= len(s.currentFrame) { + s.dequeueNextFrame() + } + } +} + +func (s *ReceiveStream) dequeueNextFrame() { + var offset protocol.ByteCount + // We're done with the last frame. Release the buffer. + if s.currentFrameDone != nil { + s.currentFrameDone() + } + offset, s.currentFrame, s.currentFrameDone = s.frameQueue.Pop() + s.currentFrameIsLast = offset+protocol.ByteCount(len(s.currentFrame)) >= s.finalOffset && !s.cancelledRemotely + s.readPosInFrame = 0 +} + +// CancelRead aborts receiving on this stream. +// It instructs the peer to stop transmitting stream data. +// Read will unblock immediately, and future Read calls will fail. +// When called multiple times or after reading the io.EOF it is a no-op. +func (s *ReceiveStream) CancelRead(errorCode StreamErrorCode) { + s.mutex.Lock() + queuedNewControlFrame := s.cancelReadImpl(errorCode) + completed := s.isNewlyCompleted() + s.mutex.Unlock() + + if queuedNewControlFrame { + s.sender.onHasStreamControlFrame(s.streamID, s) + } + if completed { + s.flowController.Abandon() + s.sender.onStreamCompleted(s.streamID) + } +} + +func (s *ReceiveStream) cancelReadImpl(errorCode qerr.StreamErrorCode) (queuedNewControlFrame bool) { + if s.cancelledLocally { // duplicate call to CancelRead + return false + } + if s.closeForShutdownErr != nil { + return false + } + s.cancelledLocally = true + if s.errorRead || s.cancelledRemotely { + return false + } + s.queuedStopSending = true + s.cancelErr = &StreamError{StreamID: s.streamID, ErrorCode: errorCode, Remote: false} + s.signalRead() + return true +} + +func (s *ReceiveStream) handleStreamFrame(frame *wire.StreamFrame, now monotime.Time) error { + s.mutex.Lock() + err := s.handleStreamFrameImpl(frame, now) + completed := s.isNewlyCompleted() + size, callback := s.takeReceiveFinalSizeCallback() + s.mutex.Unlock() + + if completed { + s.flowController.Abandon() + s.sender.onStreamCompleted(s.streamID) + } + if callback != nil { + callback(size) + } + return err +} + +func (s *ReceiveStream) handleStreamFrameImpl(frame *wire.StreamFrame, now monotime.Time) error { + if s.closeForShutdownErr != nil { + return nil + } + maxOffset := frame.Offset + frame.DataLen() + if err := s.flowController.UpdateHighestReceived(maxOffset, frame.Fin, now); err != nil { + return err + } + if frame.Fin { + s.finalOffset = maxOffset + } + if s.cancelledLocally { + return nil + } + if err := s.frameQueue.Push(frame.Data, frame.Offset, frame.PutBack); err != nil { + return err + } + s.signalRead() + return nil +} + +func (s *ReceiveStream) handleResetStreamFrame(frame *wire.ResetStreamFrame, now monotime.Time) error { + s.mutex.Lock() + err := s.handleResetStreamFrameImpl(frame, now) + completed := s.isNewlyCompleted() + size, callback := s.takeReceiveFinalSizeCallback() + s.mutex.Unlock() + + if completed { + s.sender.onStreamCompleted(s.streamID) + } + if callback != nil { + callback(size) + } + return err +} + +func (s *ReceiveStream) handleResetStreamFrameImpl(frame *wire.ResetStreamFrame, now monotime.Time) error { + if s.closeForShutdownErr != nil { + return nil + } + if err := s.flowController.UpdateHighestReceived(frame.FinalSize, true, now); err != nil { + return err + } + s.finalOffset = frame.FinalSize + + // senders are allowed to reduce the reliable size, but frames might have been reordered + if (!s.cancelledRemotely && s.reliableSize == 0) || frame.ReliableSize < s.reliableSize { + s.reliableSize = frame.ReliableSize + } + if s.readPos >= s.reliableSize { + // calling Abandon multiple times is a no-op + s.flowController.Abandon() + } + // ignore duplicate RESET_STREAM frames for this stream (after checking their final offset) + if s.cancelledRemotely { + return nil + } + + // don't save the error if the RESET_STREAM frames was received after CancelRead was called + if s.cancelledLocally { + return nil + } + s.cancelledRemotely = true + s.cancelErr = &StreamError{StreamID: s.streamID, ErrorCode: frame.ErrorCode, Remote: true} + s.signalRead() + return nil +} + +func (s *ReceiveStream) takeReceiveFinalSizeCallback() (int64, func(int64)) { + if s.finalOffset == protocol.MaxByteCount { + return 0, nil + } + callback := s.receiveFinalSizeCallback + s.receiveFinalSizeCallback = nil + return int64(s.finalOffset), callback +} + +func (s *ReceiveStream) getControlFrame(now monotime.Time) (_ ackhandler.Frame, ok, hasMore bool) { + s.mutex.Lock() + defer s.mutex.Unlock() + + if !s.queuedStopSending && !s.queuedMaxStreamData { + return ackhandler.Frame{}, false, false + } + if s.queuedStopSending { + s.queuedStopSending = false + return ackhandler.Frame{ + Frame: &wire.StopSendingFrame{StreamID: s.streamID, ErrorCode: s.cancelErr.ErrorCode}, + }, true, s.queuedMaxStreamData + } + + s.queuedMaxStreamData = false + return ackhandler.Frame{ + Frame: &wire.MaxStreamDataFrame{ + StreamID: s.streamID, + MaximumStreamData: s.flowController.GetWindowUpdate(now), + }, + }, true, false +} + +// SetReadDeadline sets the deadline for future Read calls and +// any currently-blocked Read call. +// A zero value for t means Read will not time out. +func (s *ReceiveStream) SetReadDeadline(t time.Time) error { + s.mutex.Lock() + s.deadline = monotime.FromTime(t) + s.mutex.Unlock() + s.signalRead() + return nil +} + +// CloseForShutdown closes a stream abruptly. +// It makes Read unblock (and return the error) immediately. +// The peer will NOT be informed about this: the stream is closed without sending a FIN or RESET. +func (s *ReceiveStream) closeForShutdown(err error) { + s.mutex.Lock() + s.closeForShutdownErr = err + s.receiveFinalSizeCallback = nil + s.mutex.Unlock() + s.signalRead() +} + +// signalRead performs a non-blocking send on the readChan +func (s *ReceiveStream) signalRead() { + select { + case s.readChan <- struct{}{}: + default: + } +} diff --git a/third_party/quic-go/receive_stream_test.go b/third_party/quic-go/receive_stream_test.go new file mode 100644 index 0000000..f0f3c87 --- /dev/null +++ b/third_party/quic-go/receive_stream_test.go @@ -0,0 +1,1140 @@ +package quic + +import ( + "fmt" + "io" + "os" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type readerWithTimeout struct { + io.Reader + Timeout time.Duration +} + +func (r *readerWithTimeout) Read(p []byte) (n int, err error) { + done := make(chan struct{}) + go func() { + defer close(done) + n, err = r.Reader.Read(p) + }() + + select { + case <-done: + return n, err + case <-time.After(r.Timeout): + return 0, fmt.Errorf("read timeout after %s", r.Timeout) + } +} + +type peeker interface { + Peek(b []byte) (int, error) +} + +type peekerWithTimeout struct { + Peeker peeker + Timeout time.Duration +} + +func (p *peekerWithTimeout) Peek(b []byte) (n int, err error) { + done := make(chan struct{}) + go func() { + defer close(done) + n, err = p.Peeker.Peek(b) + }() + + select { + case <-done: + return n, err + case <-time.After(p.Timeout): + return 0, fmt.Errorf("peek timeout after %s", p.Timeout) + } +} + +func TestReceiveStreamReadData(t *testing.T) { + mockFC := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, mockFC) + + // read an entire frame + now := monotime.Now() + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte{0xde, 0xad, 0xbe, 0xef}}, now)) + b := make([]byte, 4) + n, err := (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 4, n) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, b) + + // split a frame across multiple reads + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 4, Data: []byte{0xca, 0xfe, 0xba, 0xbe}}, now)) + b = make([]byte, 2) + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, []byte{0xca, 0xfe}, b) + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, []byte{0xba, 0xbe}, b) + + // combine two frames + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 8, Data: []byte{'f', 'o', 'o'}}, now)) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 11, Data: []byte{'b', 'a', 'r'}}, now)) + b = make([]byte, 6) + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 6, n) + require.Equal(t, []byte{'f', 'o', 'o', 'b', 'a', 'r'}, b) + + // reordered frames + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 17, Data: []byte{'b', 'a', 'z'}}, now)) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 14, Data: []byte{'f', 'o', 'o'}}, now)) + b = make([]byte, 6) + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 6, n) + require.Equal(t, []byte{'f', 'o', 'o', 'b', 'a', 'z'}, b) +} + +func TestReceiveStreamPeekData(t *testing.T) { + mockFC := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, mockFC) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foo")}, monotime.Now())) + b := make([]byte, 2) + n, err := (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, []byte("fo"), b) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("bar"), Offset: 3}, monotime.Now())) + b = make([]byte, 6) + n, err = (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.NoError(t, err) + require.Equal(t, 6, n) + require.Equal(t, []byte("foobar"), b) + + _, err = str.Read([]byte{0, 0}) + require.NoError(t, err) + + b = make([]byte, 2) + n, err = (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, []byte("ob"), b) + b = make([]byte, 4) + n, err = (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.NoError(t, err) + require.Equal(t, 4, n) + require.Equal(t, []byte("obar"), b) +} + +func TestReceiveStreamBlockRead(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + errChan := make(chan error, 1) + start := monotime.Now() + go func() { + frame := &wire.StreamFrame{Data: []byte{0xde, 0xad}} + time.Sleep(time.Hour) + errChan <- str.handleStreamFrame(frame, monotime.Now()) + }() + + n, err := (&readerWithTimeout{Reader: str, Timeout: 2 * time.Hour}).Read(make([]byte, 2)) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, time.Hour, monotime.Since(start)) + require.NoError(t, <-errChan) + }) +} + +func TestReceiveStreamBlockPeek(t *testing.T) { + t.Run("single STREAM frame", func(t *testing.T) { + testReceiveStreamBlockPeek(t, false) + }) + + t.Run("multiple STREAM frames", func(t *testing.T) { + testReceiveStreamBlockPeek(t, true) + }) +} + +func testReceiveStreamBlockPeek(t *testing.T, smallWrites bool) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + errChan := make(chan error, 2) + start := monotime.Now() + go func() { + if smallWrites { + time.Sleep(30 * time.Minute) + errChan <- str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foo")}, monotime.Now()) + time.Sleep(30 * time.Minute) + errChan <- str.handleStreamFrame(&wire.StreamFrame{Offset: 3, Data: []byte("bar")}, monotime.Now()) + } else { + time.Sleep(time.Hour) + errChan <- str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now()) + } + }() + + b := make([]byte, 6) + n, err := (&peekerWithTimeout{Peeker: str, Timeout: 2 * time.Hour}).Peek(b) + require.NoError(t, err) + require.Equal(t, 6, n) + require.Equal(t, []byte("foobar"), b) + require.Equal(t, time.Hour, monotime.Since(start)) + require.NoError(t, <-errChan) + if smallWrites { + require.NoError(t, <-errChan) + } + }) +} + +func TestReceiveStreamReadOverlappingData(t *testing.T) { + mockFC := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, mockFC) + + // receive the same frame multiple times + now := monotime.Now() + for range 3 { + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte{0xde, 0xad, 0xbe, 0xef}}, now)) + } + b := make([]byte, 4) + n, err := (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 4, n) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, b) + + // receive overlapping data + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 4, Data: []byte("foob")}, now)) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 6, Data: []byte("obar")}, now)) + b = make([]byte, 6) + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(b) + require.NoError(t, err) + require.Equal(t, 6, n) + require.Equal(t, []byte("foobar"), b) +} + +func TestReceiveStreamFlowControlUpdates(t *testing.T) { + t.Run("stream", func(t *testing.T) { + testReceiveStreamFlowControlUpdates(t, true, false) + }) + + t.Run("connection", func(t *testing.T) { + testReceiveStreamFlowControlUpdates(t, false, true) + }) +} + +func testReceiveStreamFlowControlUpdates(t *testing.T, hasStreamWindowUpdate, hasConnWindowUpdate bool) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + streamReceiveWindow := protocol.MaxByteCount + connReceiveWindow := protocol.MaxByteCount + if hasStreamWindowUpdate { + streamReceiveWindow = 4 + } + if hasConnWindowUpdate { + connReceiveWindow = 4 + } + mockFC := newTestStreamFlowControllerWithWindows(42, 0, streamReceiveWindow, connReceiveWindow) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(streamID, mockSender, mockFC) + + now := monotime.Now() + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte{0xde, 0xad, 0xbe, 0xef}}, now)) + + if hasStreamWindowUpdate { + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + } + if hasConnWindowUpdate { + mockSender.EXPECT().onHasConnectionData() + } + n, err := (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(make([]byte, 3)) + require.NoError(t, err) + require.Equal(t, 3, n) + require.True(t, mockCtrl.Satisfied()) + + if hasStreamWindowUpdate { + now = now.Add(time.Second) + f, ok, hasMore := str.getControlFrame(now) + require.True(t, ok) + require.Equal(t, &wire.MaxStreamDataFrame{StreamID: streamID, MaximumStreamData: protocol.ByteCount(n) + streamReceiveWindow}, f.Frame) + require.False(t, hasMore) + } + if hasConnWindowUpdate { + _, ok, hasMore := str.getControlFrame(now) + require.False(t, ok) + require.False(t, hasMore) + } +} + +func TestReceiveStreamDeadlineInThePast(t *testing.T) { + t.Run("read", func(t *testing.T) { + testReceiveStreamDeadlineInThePast(t, true, func(str *ReceiveStream, b []byte) (int, error) { + return str.Read(b) + }) + }) + t.Run("peek", func(t *testing.T) { + testReceiveStreamDeadlineInThePast(t, false, func(str *ReceiveStream, b []byte) (int, error) { + return str.Peek(b) + }) + }) +} + +func testReceiveStreamDeadlineInThePast(t *testing.T, consumesBytes bool, op func(*ReceiveStream, []byte) (int, error)) { + mockFC := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, mockFC) + + // no data is read when the deadline is in the past + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now())) + require.NoError(t, str.SetReadDeadline(time.Now().Add(-time.Second))) + b := make([]byte, 6) + n, err := op(str, b) + require.Error(t, err) + require.Zero(t, n) + require.ErrorIs(t, err, errDeadline) + + // data is read when the deadline is in the future + require.NoError(t, str.SetReadDeadline(time.Now().Add(time.Second))) + n, err = op(str, b) + require.NoError(t, err) + require.Equal(t, 6, n) +} + +func TestReceiveStreamDeadlineRemoval(t *testing.T) { + t.Run("read", func(t *testing.T) { + testReceiveStreamDeadlineRemoval(t, func(str *ReceiveStream) error { + _, err := str.Read([]byte{0}) + return err + }) + }) + t.Run("peek", func(t *testing.T) { + testReceiveStreamDeadlineRemoval(t, func(str *ReceiveStream) error { + _, err := str.Peek([]byte{0}) + return err + }) + }) +} + +func testReceiveStreamDeadlineRemoval(t *testing.T, op func(*ReceiveStream) error) { + synctest.Test(t, func(t *testing.T) { + mockFC := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, mockFC) + + const deadline = time.Minute + require.NoError(t, str.SetReadDeadline(time.Now().Add(deadline))) + errChan := make(chan error, 1) + go func() { + errChan <- op(str) + }() + select { + case err := <-errChan: + t.Fatalf("should not have returned yet: %v", err) + case <-time.After(deadline / 2): + } + + // remove the deadline after a while (but before it expires) + require.NoError(t, str.SetReadDeadline(time.Time{})) + + // no deadline set: should not return at all + select { + case err := <-errChan: + t.Fatalf("should not have returned yet: %v", err) + case <-time.After(2 * deadline): + } + + // now set the deadline to the past to make it return immediately + require.NoError(t, str.SetReadDeadline(time.Now().Add(-time.Second))) + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + default: + t.Fatal("timeout") + } + }) +} + +func TestReceiveStreamDeadlineExtension(t *testing.T) { + t.Run("read", func(t *testing.T) { + testReceiveStreamDeadlineExtension(t, func(str *ReceiveStream) error { + _, err := str.Read([]byte{0}) + return err + }) + }) + t.Run("peek", func(t *testing.T) { + testReceiveStreamDeadlineExtension(t, func(str *ReceiveStream) error { + _, err := str.Peek([]byte{0}) + return err + }) + }) +} + +func testReceiveStreamDeadlineExtension(t *testing.T, op func(*ReceiveStream) error) { + synctest.Test(t, func(t *testing.T) { + mockFC := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, mockFC) + + start := monotime.Now() + deadline := 5 * time.Second + require.NoError(t, str.SetReadDeadline(time.Now().Add(deadline))) + errChan := make(chan error, 1) + go func() { + errChan <- op(str) + }() + select { + case err := <-errChan: + t.Fatalf("should not have returned yet: %v", err) + case <-time.After(deadline / 2): + } + + // extend the deadline + require.NoError(t, str.SetReadDeadline(time.Now().Add(deadline))) + select { + case err := <-errChan: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Equal(t, start.Add(deadline*3/2), monotime.Now()) + case <-time.After(deadline + time.Nanosecond): + t.Fatal("timeout") + } + }) +} + +func TestReceiveStreamEOFWithData(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + now := monotime.Now() + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 2, Data: []byte{0xbe, 0xef}, Fin: true}, now)) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte{0xde, 0xad}}, now)) + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + + // peeking doesn't return an EOF + b := make([]byte, 4) + n, err := (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.NoError(t, err) + require.Equal(t, 4, n) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, b) + + // peeking returns the EOF, if more data is being peeked + b = make([]byte, 6) + n, err = (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.ErrorIs(t, err, io.EOF) + require.Equal(t, 4, n) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, b[:n]) + + // reading returns the EOF + strWithTimeout := &readerWithTimeout{Reader: str, Timeout: time.Second} + b = make([]byte, 6) + n, err = strWithTimeout.Read(b) + require.ErrorIs(t, err, io.EOF) + require.Equal(t, 4, n) + require.Equal(t, []byte{0xde, 0xad, 0xbe, 0xef}, b[:n]) + n, err = strWithTimeout.Read(b) + require.Zero(t, n) + require.ErrorIs(t, err, io.EOF) +} + +func TestReceiveStreamPeekEOF(t *testing.T) { + t.Run("long peek", func(t *testing.T) { + testReceiveStreamPeekEOF(t, true) + }) + t.Run("exact peek", func(t *testing.T) { + testReceiveStreamPeekEOF(t, false) + }) +} + +func testReceiveStreamPeekEOF(t *testing.T, longPeek bool) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Offset: 3, Data: []byte("bar"), Fin: true}, monotime.Now())) + + type result struct { + err error + data []byte + } + resultChan := make(chan result, 1) + go func() { + b := make([]byte, 6) + if longPeek { + b = make([]byte, 8) + } + n, err := (&peekerWithTimeout{Peeker: str, Timeout: time.Hour}).Peek(b) + resultChan <- result{err: err, data: b[:n]} + }() + + synctest.Wait() + + select { + case result := <-resultChan: + t.Fatalf("peek should not have returned yet: %v", result.err) + default: + } + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("f")}, monotime.Now())) + + synctest.Wait() + + select { + case result := <-resultChan: + t.Fatalf("peek should not have returned yet: %v", result.err) + default: + } + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("oo"), Offset: 1}, monotime.Now())) + + synctest.Wait() + + select { + case result := <-resultChan: + if longPeek { + assert.ErrorIs(t, result.err, io.EOF) + } else { + assert.NoError(t, result.err) + } + require.Equal(t, []byte("foobar"), result.data) + default: + t.Fatal("peek should have returned") + } + }) +} + +func TestReceiveStreamImmediateFINs(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Fin: true}, monotime.Now())) + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + + // peeking returns the EOF + n, err := (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(make([]byte, 4)) + require.ErrorIs(t, err, io.EOF) + require.Equal(t, 0, n) + + // and so does reading + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(make([]byte, 4)) + require.Zero(t, n) + require.ErrorIs(t, err, io.EOF) +} + +func TestReceiveStreamFinalSizeCallbackAfterFIN(t *testing.T) { + fc := newTestStreamFlowController(42) + str := newReceiveStream(42, nil, fc) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, monotime.Now())) + + var size int64 + str.SetReceiveFinalSizeCallback(func(s int64) { size = s }) + require.EqualValues(t, 6, size) + + str.closeForShutdown(assert.AnError) + str.SetReceiveFinalSizeCallback(func(s int64) { size = s }) + require.EqualValues(t, 6, size) +} + +func TestReceiveStreamFinalSizeCallbackAfterCancelRead(t *testing.T) { + mockCtrl := gomock.NewController(t) + fc := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, fc) + + var size int64 + var called bool + str.SetReceiveFinalSizeCallback(func(s int64) { + size, called = s, true + }) + require.False(t, called) + + mockSender.EXPECT().onHasStreamControlFrame(str.StreamID(), str) + str.CancelRead(1234) + require.False(t, called) + + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + require.NoError(t, str.handleResetStreamFrame( + &wire.ResetStreamFrame{StreamID: 42, ErrorCode: 4321, FinalSize: 42}, + monotime.Now(), + )) + + require.True(t, called) + require.EqualValues(t, 42, size) +} + +func TestReceiveStreamFinalSizeCallbackRemoved(t *testing.T) { + str := newReceiveStream(42, nil, newTestStreamFlowController(42)) + str.SetReceiveFinalSizeCallback(func(int64) { t.Fatal("callback called") }) + str.SetReceiveFinalSizeCallback(nil) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Fin: true}, monotime.Now())) +} + +func TestReceiveStreamCloseForShutdown(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + strWithTimeout := &readerWithTimeout{Reader: str, Timeout: time.Minute} + + // Test immediate return of reads + readErrChan := make(chan error, 1) + peekErrChan := make(chan error, 1) + go func() { + _, err := strWithTimeout.Read([]byte{0}) + readErrChan <- err + }() + go func() { + _, err := (&peekerWithTimeout{Peeker: str, Timeout: time.Minute}).Peek([]byte{0}) + peekErrChan <- err + }() + + synctest.Wait() + + select { + case err := <-readErrChan: + t.Fatalf("read returned before closeForShutdown: %v", err) + case err := <-peekErrChan: + t.Fatalf("peek returned before closeForShutdown: %v", err) + default: + } + + str.closeForShutdown(assert.AnError) + synctest.Wait() + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Fin: true}, monotime.Now())) + str.SetReceiveFinalSizeCallback(func(int64) { t.Fatal("callback called") }) + + select { + case err := <-readErrChan: + require.ErrorIs(t, err, assert.AnError) + default: + t.Fatal("read should have returned") + } + select { + case err := <-peekErrChan: + require.ErrorIs(t, err, assert.AnError) + default: + t.Fatal("peek should have returned") + } + + // following calls to Peek should return the error + n, err := (&peekerWithTimeout{Peeker: str, Timeout: time.Minute}).Peek([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, assert.AnError) + + // following calls to Read should return the error + n, err = strWithTimeout.Read([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, assert.AnError) + + // receiving a RESET_STREAM frame after closeForShutdown does nothing + require.NoError(t, str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1234, FinalSize: 42}, monotime.Now())) + n, err = strWithTimeout.Read([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, assert.AnError) + + // calling CancelRead after closeForShutdown does nothing + str.CancelRead(1234) + n, err = strWithTimeout.Read([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, assert.AnError) + }) +} + +func TestReceiveStreamCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + strWithTimeout := &readerWithTimeout{Reader: str, Timeout: 2 * time.Second} + + mockSender.EXPECT().onHasStreamControlFrame(str.StreamID(), gomock.Any()) + readErrChan := make(chan error, 1) + peekErrChan := make(chan error, 1) + go func() { + _, err := strWithTimeout.Read([]byte{0}) + readErrChan <- err + }() + go func() { + _, err := (&peekerWithTimeout{Peeker: str, Timeout: 2 * time.Second}).Peek([]byte{0}) + peekErrChan <- err + }() + + synctest.Wait() + + str.CancelRead(1234) + // this queues a STOP_SENDING frame + f, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.StopSendingFrame{StreamID: 42, ErrorCode: 1234}, f.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + synctest.Wait() + + select { + case err := <-readErrChan: + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: false}) + default: + t.Fatal("Read was not unblocked") + } + select { + case err := <-peekErrChan: + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: false}) + default: + t.Fatal("Peek was not unblocked") + } + + // further calls to Peek return the error + n, err := (&peekerWithTimeout{Peeker: str, Timeout: 2 * time.Second}).Peek([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: false}) + + // further Read calls return the error + n, err = strWithTimeout.Read([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: false}) + + // calling CancelRead again does nothing + // especially: + // 1. no more calls to onHasStreamControlFrame + // 2. no changes of the error code returned by Read + str.CancelRead(1234) + str.CancelRead(4321) + n, err = strWithTimeout.Read([]byte{0}) + require.Zero(t, n) + // error code unchanged + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: false}) + require.True(t, mockCtrl.Satisfied()) + + // receiving the FIN bit has no effect + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + // receive two of them, to make sure onStreamCompleted is not called twice + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, monotime.Now())) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, monotime.Now())) + require.True(t, mockCtrl.Satisfied()) + + // receiving a RESET_STREAM frame after CancelRead has no effect + require.NoError(t, str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 4321, FinalSize: 6}, monotime.Now())) + n, err = strWithTimeout.Read([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: false}) + }) +} + +func TestReceiveStreamCancelReadAbandonsUnreadData(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + connFC := newConnectionFlowController(6, protocol.MaxByteCount, nil, utils.NewRTTStats(), utils.DefaultLogger) + fc := newStreamFlowController( + streamID, + connFC, + protocol.MaxByteCount, + protocol.MaxByteCount, + 0, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(streamID, mockSender, fc) + + now := monotime.Now() + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, now)) + require.Zero(t, connFC.GetWindowUpdate(now)) + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + mockSender.EXPECT().onStreamCompleted(streamID) + str.CancelRead(1337) + require.True(t, mockCtrl.Satisfied()) + + // CancelRead completes the receive side once the final offset is known. + // Completion abandons the unread stream data, which returns connection-level + // flow-control credit and makes a connection window update available. + require.NotZero(t, connFC.GetWindowUpdate(now)) +} + +func TestReceiveStreamCancelReadAfterFIN(t *testing.T) { + t.Run("FIN not read", func(t *testing.T) { + testReceiveStreamCancelReadAfterFIN(t, false) + }) + t.Run("FIN read", func(t *testing.T) { + testReceiveStreamCancelReadAfterFIN(t, true) + }) +} + +func testReceiveStreamCancelReadAfterFIN(t *testing.T, finRead bool) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, monotime.Now())) + if finRead { + n, err := str.Read(make([]byte, 10)) + require.ErrorIs(t, err, io.EOF) + require.Equal(t, 6, n) + } + + // if the FIN was received, but not read yet, a STOP_SENDING frame is queued + if !finRead { + mockSender.EXPECT().onHasStreamControlFrame(str.StreamID(), str) + } + str.CancelRead(1337) + f, ok, hasMore := str.getControlFrame(monotime.Now()) + // if the EOF was already read, no STOP_SENDING frame is queued + if finRead { + require.False(t, ok) + require.False(t, hasMore) + } else { + require.True(t, ok) + require.Equal(t, &wire.StopSendingFrame{StreamID: 42, ErrorCode: 1337}, f.Frame) + require.False(t, hasMore) + } + + // Read returns the error... + n, err := str.Read([]byte{0}) + require.Zero(t, n) + // ... and Peek returns the same error + n, peekErr := (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek([]byte{0}) + require.Zero(t, n) + if finRead { + assert.ErrorIs(t, err, io.EOF) + assert.ErrorIs(t, peekErr, io.EOF) + } else { + assert.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: false}) + assert.ErrorIs(t, peekErr, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: false}) + } +} + +func TestReceiveStreamResetAbandonsUnreadData(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + connFC := newConnectionFlowController(6, protocol.MaxByteCount, nil, utils.NewRTTStats(), utils.DefaultLogger) + fc := newStreamFlowController( + streamID, + connFC, + protocol.MaxByteCount, + protocol.MaxByteCount, + 0, + utils.NewRTTStats(), + utils.DefaultLogger, + ) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(streamID, mockSender, fc) + + now := monotime.Now() + require.Zero(t, connFC.GetWindowUpdate(now)) + require.NoError(t, str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: streamID, ErrorCode: 1234, FinalSize: 6}, now)) + + // A RESET_STREAM with no reliable data abandons the unread final size. + // That returns connection-level flow-control credit and makes a connection + // window update available. + require.NotZero(t, connFC.GetWindowUpdate(now)) +} + +func TestReceiveStreamReset(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + strWithTimeout := &readerWithTimeout{Reader: str, Timeout: 2 * time.Second} + + readErrChan := make(chan error, 1) + peekErrChan := make(chan error, 1) + go func() { + _, err := strWithTimeout.Read([]byte{0}) + readErrChan <- err + }() + go func() { + _, err := (&peekerWithTimeout{Peeker: str, Timeout: 2 * time.Second}).Peek([]byte{0}) + peekErrChan <- err + }() + + synctest.Wait() + + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + require.NoError(t, str.handleResetStreamFrame( + &wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1234, FinalSize: 42}, + monotime.Now(), + )) + + synctest.Wait() + + select { + case err := <-readErrChan: + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: true}) + default: + t.Fatal("Read was not unblocked") + } + select { + case err := <-peekErrChan: + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: true}) + default: + t.Fatal("Peek was not unblocked") + } + + // further calls to Peek return the error + n, err := (&peekerWithTimeout{Peeker: str, Timeout: 2 * time.Second}).Peek([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: true}) + + // further calls to Read return the error + _, err = strWithTimeout.Read([]byte{0}) + require.Equal(t, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: true}, err) + + // further RESET_STREAM frames have no effect + require.NoError(t, str.handleResetStreamFrame( + &wire.ResetStreamFrame{StreamID: 42, ErrorCode: 4321, FinalSize: 42}, + monotime.Now(), + )) + + n, err = str.Read([]byte{0}) + require.Zero(t, n) + // error code unchanged + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: true}) + + // CancelRead after a RESET_STREAM frame has no effect + str.CancelRead(100) + n, err = str.Read([]byte{0}) + require.Zero(t, n) + // error code and remote flag unchanged + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1234, Remote: true}) + }) +} + +func TestReceiveStreamResetAfterFINRead(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + require.NoError(t, str.handleStreamFrame( + &wire.StreamFrame{StreamID: 42, Data: []byte("foobar"), Fin: true}, + monotime.Now(), + )) + n, err := str.Read(make([]byte, 6)) + require.Equal(t, 6, n) + require.ErrorIs(t, err, io.EOF) + // make sure that onStreamCompleted was called due to the EOF + require.True(t, mockCtrl.Satisfied()) + + // Now receive a RESET_STREAM frame. + // We don't expect any more calls to onStreamCompleted. + require.NoError(t, str.handleResetStreamFrame( + &wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1234, FinalSize: 6}, + monotime.Now(), + )) + // now read the error + n, err = str.Read([]byte{0}) + require.Error(t, err) + require.Zero(t, n) +} + +// Calling Read concurrently doesn't make any sense (and is forbidden), +// but we still want to make sure that we don't complete the stream more than once +// if the user misuses our API. +// This would lead to an INTERNAL_ERROR ("tried to delete unknown outgoing stream"), +// which can be hard to debug. +// Note that even without the protection built into the receiveStream, this test +// is very timing-dependent, and would need to run a few hundred times to trigger the failure. +func TestReceiveStreamConcurrentReads(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + var numCompleted atomic.Int32 + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)).Do(func(protocol.StreamID) { + numCompleted.Add(1) + }).AnyTimes() + + const num = 3 + resultChan := make(chan struct { + n int + err error + }, num) + for range num { + go func() { + n, err := str.Read(make([]byte, 8)) + resultChan <- struct { + n int + err error + }{n: n, err: err} + }() + } + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, monotime.Now())) + synctest.Wait() + + var bytesRead int + for range num { + select { + case res := <-resultChan: + bytesRead += res.n + require.ErrorIs(t, res.err, io.EOF) + default: + t.Fatal("read should have returned") + } + } + require.Equal(t, 6, bytesRead) + require.Equal(t, int32(1), numCompleted.Load()) + }) +} + +func TestReceiveStreamResetStreamAtBeforeReadOffset(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now())) + b := make([]byte, 3) + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("foo"), b) + + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10, ReliableSize: 3}, monotime.Now()) + require.True(t, mockCtrl.Satisfied()) + + // Peek returns the error + n, err = str.Peek([]byte{0}) + require.Zero(t, n) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: true}) + + // Read returns the error + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + n, err = str.Read([]byte{0}) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: true}) + require.Zero(t, n) +} + +func TestReceiveStreamResetStreamAtAfterReadOffset(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now())) + b := make([]byte, 2) + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, []byte("fo"), b) + + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10, ReliableSize: 6}, monotime.Now()) + require.True(t, mockCtrl.Satisfied()) + + // Peek returns no error when peeking up to the reliable size... + b = make([]byte, 4) + n, err = (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.NoError(t, err) + require.Equal(t, 4, n) + require.Equal(t, []byte("obar"), b) + + // ... but returns the error when peeking beyond the reliable size + b = make([]byte, 5) + n, err = (&peekerWithTimeout{Peeker: str, Timeout: time.Second}).Peek(b) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: true}) + require.Equal(t, 4, n) + require.Equal(t, []byte("obar"), b[:n]) + + // Read returns the error after reading up to the reliable size + b = make([]byte, 2) + n, err = str.Read(b) + require.NoError(t, err) + require.Equal(t, 2, n) + require.Equal(t, []byte("ob"), b) + require.True(t, mockCtrl.Satisfied()) + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + n, err = str.Read(b) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: true}) + require.Equal(t, 2, n) + require.Equal(t, []byte("ar"), b) +} + +func TestReceiveStreamMultipleResetStreamAt(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now())) + + b := make([]byte, 3) + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("foo"), b) + require.True(t, mockCtrl.Satisfied()) + + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10, ReliableSize: 6}, monotime.Now()) + require.True(t, mockCtrl.Satisfied()) + + // receiving a reordered RESET_STREAM_AT frame has no effect + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10, ReliableSize: 8}, monotime.Now()) + require.True(t, mockCtrl.Satisfied()) + + // receiving a RESET_STREAM_AT frame with a smaller reliable size is valid + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10, ReliableSize: 3}, monotime.Now()) + + // Read returns the error + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + n, err = str.Read(b) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: true}) + require.Zero(t, n) +} + +func TestReceiveStreamResetStreamAtAfterResetStream(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowController(42) + mockSender := NewMockStreamSender(mockCtrl) + str := newReceiveStream(42, mockSender, mockFC) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now())) + + b := make([]byte, 3) + n, err := str.Read(b) + require.NoError(t, err) + require.Equal(t, 3, n) + require.Equal(t, []byte("foo"), b) + require.True(t, mockCtrl.Satisfied()) + + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10}, monotime.Now()) + require.True(t, mockCtrl.Satisfied()) + + // receiving a reordered RESET_STREAM_AT frame has no effect + str.handleResetStreamFrame(&wire.ResetStreamFrame{StreamID: 42, ErrorCode: 1337, FinalSize: 10, ReliableSize: 8}, monotime.Now()) + require.True(t, mockCtrl.Satisfied()) + + // Read returns the error + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(42)) + n, err = str.Read(b) + require.ErrorIs(t, err, &StreamError{StreamID: 42, ErrorCode: 1337, Remote: true}) + require.Zero(t, n) +} diff --git a/third_party/quic-go/retransmission_queue.go b/third_party/quic-go/retransmission_queue.go new file mode 100644 index 0000000..5120518 --- /dev/null +++ b/third_party/quic-go/retransmission_queue.go @@ -0,0 +1,158 @@ +package quic + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/ackhandler" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +type framesToRetransmit struct { + crypto []*wire.CryptoFrame + other []wire.Frame +} + +type retransmissionQueue struct { + initial *framesToRetransmit + handshake *framesToRetransmit + appData framesToRetransmit +} + +func newRetransmissionQueue() *retransmissionQueue { + return &retransmissionQueue{ + initial: &framesToRetransmit{}, + handshake: &framesToRetransmit{}, + } +} + +func (q *retransmissionQueue) addInitial(f wire.Frame) { + if q.initial == nil { + return + } + if cf, ok := f.(*wire.CryptoFrame); ok { + q.initial.crypto = append(q.initial.crypto, cf) + return + } + q.initial.other = append(q.initial.other, f) +} + +func (q *retransmissionQueue) addHandshake(f wire.Frame) { + if q.handshake == nil { + return + } + if cf, ok := f.(*wire.CryptoFrame); ok { + q.handshake.crypto = append(q.handshake.crypto, cf) + return + } + q.handshake.other = append(q.handshake.other, f) +} + +func (q *retransmissionQueue) addAppData(f wire.Frame) { + switch f := f.(type) { + case *wire.StreamFrame: + panic("STREAM frames are handled with their respective streams.") + case *wire.CryptoFrame: + q.appData.crypto = append(q.appData.crypto, f) + default: + q.appData.other = append(q.appData.other, f) + } +} + +func (q *retransmissionQueue) HasData(encLevel protocol.EncryptionLevel) bool { + //nolint:exhaustive // 0-RTT data is retransmitted in 1-RTT packets. + switch encLevel { + case protocol.EncryptionInitial: + return q.initial != nil && + (len(q.initial.crypto) > 0 || len(q.initial.other) > 0) + case protocol.EncryptionHandshake: + return q.handshake != nil && + (len(q.handshake.crypto) > 0 || len(q.handshake.other) > 0) + case protocol.Encryption1RTT: + return len(q.appData.crypto) > 0 || len(q.appData.other) > 0 + } + return false +} + +func (q *retransmissionQueue) GetFrame(encLevel protocol.EncryptionLevel, maxLen protocol.ByteCount, v protocol.Version) wire.Frame { + var r *framesToRetransmit + //nolint:exhaustive // 0-RTT data is retransmitted in 1-RTT packets. + switch encLevel { + case protocol.EncryptionInitial: + r = q.initial + case protocol.EncryptionHandshake: + r = q.handshake + case protocol.Encryption1RTT: + r = &q.appData + } + if r == nil { + return nil + } + + if len(r.crypto) > 0 { + f := r.crypto[0] + newFrame, needsSplit := f.MaybeSplitOffFrame(maxLen, v) + if newFrame == nil && !needsSplit { // the whole frame fits + r.crypto = r.crypto[1:] + return f + } + if newFrame != nil { // frame was split. Leave the original frame in the queue. + return newFrame + } + } + if len(r.other) == 0 { + return nil + } + f := r.other[0] + if f.Length(v) > maxLen { + return nil + } + r.other = r.other[1:] + return f +} + +func (q *retransmissionQueue) DropPackets(encLevel protocol.EncryptionLevel) { + //nolint:exhaustive // Can only drop Initial and Handshake packet number space. + switch encLevel { + case protocol.EncryptionInitial: + q.initial = nil + case protocol.EncryptionHandshake: + q.handshake = nil + default: + panic(fmt.Sprintf("unexpected encryption level: %s", encLevel)) + } +} + +func (q *retransmissionQueue) AckHandler(encLevel protocol.EncryptionLevel) ackhandler.FrameHandler { + switch encLevel { + case protocol.EncryptionInitial: + return (*retransmissionQueueInitialAckHandler)(q) + case protocol.EncryptionHandshake: + return (*retransmissionQueueHandshakeAckHandler)(q) + case protocol.Encryption0RTT, protocol.Encryption1RTT: + return (*retransmissionQueueAppDataAckHandler)(q) + } + return nil +} + +type retransmissionQueueInitialAckHandler retransmissionQueue + +func (q *retransmissionQueueInitialAckHandler) OnAcked(wire.Frame) {} +func (q *retransmissionQueueInitialAckHandler) OnLost(f wire.Frame) { + (*retransmissionQueue)(q).addInitial(f) +} + +type retransmissionQueueHandshakeAckHandler retransmissionQueue + +func (q *retransmissionQueueHandshakeAckHandler) OnAcked(wire.Frame) {} +func (q *retransmissionQueueHandshakeAckHandler) OnLost(f wire.Frame) { + (*retransmissionQueue)(q).addHandshake(f) +} + +type retransmissionQueueAppDataAckHandler retransmissionQueue + +func (q *retransmissionQueueAppDataAckHandler) OnAcked(wire.Frame) {} +func (q *retransmissionQueueAppDataAckHandler) OnLost(f wire.Frame) { + (*retransmissionQueue)(q).addAppData(f) +} diff --git a/third_party/quic-go/retransmission_queue_test.go b/third_party/quic-go/retransmission_queue_test.go new file mode 100644 index 0000000..6117bcd --- /dev/null +++ b/third_party/quic-go/retransmission_queue_test.go @@ -0,0 +1,131 @@ +package quic + +import ( + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" +) + +func TestRetransmissionQueueFrames(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testRetransmissionQueueFrames(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testRetransmissionQueueFrames(t, protocol.EncryptionHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testRetransmissionQueueFrames(t, protocol.Encryption1RTT) + }) +} + +func testRetransmissionQueueFrames(t *testing.T, encLevel protocol.EncryptionLevel) { + q := newRetransmissionQueue() + + require.False(t, q.HasData(encLevel)) + require.Nil(t, q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1)) + + ah := q.AckHandler(encLevel) + require.NotNil(t, ah) + ah.OnLost(&wire.PingFrame{}) + require.True(t, q.HasData(encLevel)) + require.Equal(t, &wire.PingFrame{}, q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1)) + require.False(t, q.HasData(encLevel)) + require.Nil(t, q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1)) + + f := &wire.PathChallengeFrame{Data: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}} + ah.OnLost(f) + require.True(t, q.HasData(encLevel)) + require.Nil(t, q.GetFrame(encLevel, f.Length(protocol.Version1)-1, protocol.Version1)) + require.Equal(t, f, q.GetFrame(encLevel, f.Length(protocol.Version1), protocol.Version1)) + require.False(t, q.HasData(encLevel)) + + if encLevel == protocol.Encryption1RTT { + require.Panics(t, func() { ah.OnLost(&wire.StreamFrame{}) }) + } +} + +func TestRetransmissionQueueCryptoFrames(t *testing.T) { + t.Run("Initial", func(t *testing.T) { + testRetransmissionQueueCryptoFrames(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testRetransmissionQueueCryptoFrames(t, protocol.EncryptionHandshake) + }) + t.Run("1-RTT", func(t *testing.T) { + testRetransmissionQueueCryptoFrames(t, protocol.Encryption1RTT) + }) +} + +func testRetransmissionQueueCryptoFrames(t *testing.T, encLevel protocol.EncryptionLevel) { + q := newRetransmissionQueue() + + var otherEncLevel protocol.EncryptionLevel + switch encLevel { + case protocol.EncryptionInitial: + otherEncLevel = protocol.EncryptionHandshake + case protocol.EncryptionHandshake: + otherEncLevel = protocol.Encryption1RTT + case protocol.Encryption1RTT: + otherEncLevel = protocol.EncryptionInitial + } + + ah := q.AckHandler(encLevel) + require.NotNil(t, ah) + ah.OnLost(&wire.CryptoFrame{Data: []byte("foobar")}) + require.True(t, q.HasData(encLevel)) + require.False(t, q.HasData(otherEncLevel)) + require.Equal(t, &wire.CryptoFrame{Data: []byte("foobar")}, q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1)) + require.False(t, q.HasData(encLevel)) + require.Nil(t, q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1)) + + f := &wire.CryptoFrame{Offset: 100, Data: []byte("foobar")} + ah.OnLost(f) + ah.OnLost(&wire.PingFrame{}) + require.True(t, q.HasData(encLevel)) + require.False(t, q.HasData(otherEncLevel)) + // the CRYPTO frame wouldn't fit, not even if it was split + require.IsType(t, &wire.PingFrame{}, q.GetFrame(encLevel, 2, protocol.Version1)) + + f1 := q.GetFrame(encLevel, f.Length(protocol.Version1)-3, protocol.Version1) + require.NotNil(t, f1) + require.IsType(t, &wire.CryptoFrame{}, f1) + require.Equal(t, &wire.CryptoFrame{Offset: 100, Data: []byte("foo")}, f1) + f2 := q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, f2) + require.IsType(t, &wire.CryptoFrame{}, f2) + require.Equal(t, &wire.CryptoFrame{Offset: 103, Data: []byte("bar")}, f2) +} + +func TestRetransmissionQueueDropEncLevel(t *testing.T) { + q := newRetransmissionQueue() + require.Panics(t, func() { q.DropPackets(protocol.Encryption0RTT) }) + require.Panics(t, func() { q.DropPackets(protocol.Encryption1RTT) }) + + t.Run("Initial", func(t *testing.T) { + testRetransmissionQueueDropEncLevel(t, protocol.EncryptionInitial) + }) + t.Run("Handshake", func(t *testing.T) { + testRetransmissionQueueDropEncLevel(t, protocol.EncryptionHandshake) + }) +} + +func testRetransmissionQueueDropEncLevel(t *testing.T, encLevel protocol.EncryptionLevel) { + q := newRetransmissionQueue() + + ah := q.AckHandler(encLevel) + require.NotNil(t, ah) + ah.OnLost(&wire.PingFrame{}) + ah.OnLost(&wire.CryptoFrame{Data: []byte("foobar")}) + require.True(t, q.HasData(encLevel)) + q.DropPackets(encLevel) + require.False(t, q.HasData(encLevel)) + require.Nil(t, q.GetFrame(encLevel, protocol.MaxByteCount, protocol.Version1)) + + // losing more frame is a no-op + ah.OnLost(&wire.CryptoFrame{Data: []byte("foobar")}) + ah.OnLost(&wire.PingFrame{}) + require.False(t, q.HasData(encLevel)) +} diff --git a/third_party/quic-go/send_conn.go b/third_party/quic-go/send_conn.go new file mode 100644 index 0000000..75a87df --- /dev/null +++ b/third_party/quic-go/send_conn.go @@ -0,0 +1,136 @@ +package quic + +import ( + "net" + "sync/atomic" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +// A sendConn allows sending using a simple Write() on a non-connected packet conn. +type sendConn interface { + Write(b []byte, gsoSize uint16, ecn protocol.ECN) error + WriteTo([]byte, net.Addr, packetInfo) error + Close() error + LocalAddr() net.Addr + RemoteAddr() net.Addr + SetRemoteAddr(addr net.Addr) + ChangeRemoteAddr(addr net.Addr, info packetInfo) + + capabilities() connCapabilities +} + +type remoteAddrInfo struct { + addr net.Addr + oob []byte +} + +type sconn struct { + rawConn + + localAddr net.Addr + + remoteAddrInfo atomic.Pointer[remoteAddrInfo] + + logger utils.Logger + + // If GSO enabled, and we receive a GSO error for this remote address, GSO is disabled. + gotGSOError bool + // Used to catch the error sometimes returned by the first sendmsg call on Linux, + // see https://github.com/golang/go/issues/63322. + wroteFirstPacket bool +} + +var _ sendConn = &sconn{} + +func newSendConn(c rawConn, remote net.Addr, info packetInfo, logger utils.Logger) *sconn { + localAddr := c.LocalAddr() + if info.addr.IsValid() { + if udpAddr, ok := localAddr.(*net.UDPAddr); ok { + addrCopy := *udpAddr + addrCopy.IP = info.addr.AsSlice() + localAddr = &addrCopy + } + } + + oob := info.OOB() + // increase oob slice capacity, so we can add the UDP_SEGMENT and ECN control messages without allocating + l := len(oob) + oob = append(oob, make([]byte, 64)...)[:l] + sc := &sconn{ + rawConn: c, + localAddr: localAddr, + logger: logger, + } + sc.remoteAddrInfo.Store(&remoteAddrInfo{ + addr: remote, + oob: oob, + }) + return sc +} + +func (c *sconn) Write(p []byte, gsoSize uint16, ecn protocol.ECN) error { + ai := c.remoteAddrInfo.Load() + err := c.writePacket(p, ai.addr, ai.oob, gsoSize, ecn) + if err != nil && isGSOError(err) { + // disable GSO for future calls + c.gotGSOError = true + if c.logger.Debug() { + c.logger.Debugf("GSO failed when sending to %s", ai.addr) + } + // send out the packets one by one + for len(p) > 0 { + l := min(len(p), int(gsoSize)) + if err := c.writePacket(p[:l], ai.addr, ai.oob, 0, ecn); err != nil { + return err + } + p = p[l:] + } + return nil + } + return err +} + +func (c *sconn) writePacket(p []byte, addr net.Addr, oob []byte, gsoSize uint16, ecn protocol.ECN) error { + _, err := c.WritePacket(p, addr, oob, gsoSize, ecn) + if err != nil && !c.wroteFirstPacket && isPermissionError(err) { + _, err = c.WritePacket(p, addr, oob, gsoSize, ecn) + } + c.wroteFirstPacket = true + return err +} + +func (c *sconn) WriteTo(b []byte, addr net.Addr, info packetInfo) error { + _, err := c.WritePacket(b, addr, info.OOB(), 0, protocol.ECNUnsupported) + return err +} + +func (c *sconn) capabilities() connCapabilities { + capabilities := c.rawConn.capabilities() + if capabilities.GSO { + capabilities.GSO = !c.gotGSOError + } + return capabilities +} + +func (c *sconn) ChangeRemoteAddr(addr net.Addr, info packetInfo) { + c.remoteAddrInfo.Store(&remoteAddrInfo{ + addr: addr, + oob: info.OOB(), + }) +} + +func (c *sconn) SetRemoteAddr(addr net.Addr) { + for { + ai := c.remoteAddrInfo.Load() // load the current value of the pointer, so we can swap it + newAi := *ai // make a copy, so we can swap the pointer + newAi.addr = addr // update the address + if c.remoteAddrInfo.CompareAndSwap(ai, &newAi) { // if the pointer was swapped, we're done + break + } + } +} + +func (c *sconn) RemoteAddr() net.Addr { return c.remoteAddrInfo.Load().addr } +func (c *sconn) LocalAddr() net.Addr { return c.localAddr } diff --git a/third_party/quic-go/send_conn_test.go b/third_party/quic-go/send_conn_test.go new file mode 100644 index 0000000..b2e4072 --- /dev/null +++ b/third_party/quic-go/send_conn_test.go @@ -0,0 +1,135 @@ +package quic + +import ( + "net" + "net/netip" + "runtime" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +// Only if appendUDPSegmentSizeMsg actually appends a message (and isn't only a stub implementation), +// GSO is actually supported on this platform. +var platformSupportsGSO = len(appendUDPSegmentSizeMsg([]byte{}, 1337)) > 0 + +func TestSendConnLocalAndRemoteAddress(t *testing.T) { + remoteAddr := &net.UDPAddr{IP: net.IPv4(192, 168, 100, 200), Port: 1337} + rawConn := NewMockRawConn(gomock.NewController(t)) + rawConn.EXPECT().LocalAddr().Return(&net.UDPAddr{IP: net.IPv4(10, 11, 12, 13), Port: 14}).Times(2) + c := newSendConn( + rawConn, + remoteAddr, + packetInfo{addr: netip.AddrFrom4([4]byte{127, 0, 0, 42})}, + utils.DefaultLogger, + ) + require.Equal(t, "127.0.0.42:14", c.LocalAddr().String()) + require.Equal(t, remoteAddr, c.RemoteAddr()) + + // the local raw conn's local address is only used if we don't an address from the packet info + c = newSendConn(rawConn, remoteAddr, packetInfo{}, utils.DefaultLogger) + require.Equal(t, "10.11.12.13:14", c.LocalAddr().String()) +} + +func TestSendConnOOB(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("we don't OOB conn on windows, and no packet info will be available") + } + + remoteAddr := &net.UDPAddr{IP: net.IPv4(192, 168, 100, 200), Port: 1337} + rawConn := NewMockRawConn(gomock.NewController(t)) + rawConn.EXPECT().LocalAddr() + rawConn.EXPECT().capabilities().AnyTimes() + pi := packetInfo{addr: netip.IPv6Loopback()} + rawConn.EXPECT().WritePacket([]byte("foobar"), remoteAddr, pi.OOB(), uint16(0), protocol.ECT1) + require.NotEmpty(t, pi.OOB()) + c := newSendConn(rawConn, remoteAddr, pi, utils.DefaultLogger) + require.NoError(t, c.Write([]byte("foobar"), 0, protocol.ECT1)) +} + +func TestSendConnDetectGSOFailure(t *testing.T) { + if !platformSupportsGSO { + t.Skip("GSO is not supported on this platform") + } + + remoteAddr := &net.UDPAddr{IP: net.IPv4(192, 168, 100, 200), Port: 1337} + rawConn := NewMockRawConn(gomock.NewController(t)) + rawConn.EXPECT().LocalAddr() + rawConn.EXPECT().capabilities().Return(connCapabilities{GSO: true}).MinTimes(1) + c := newSendConn(rawConn, remoteAddr, packetInfo{}, utils.DefaultLogger) + gomock.InOrder( + rawConn.EXPECT().WritePacket([]byte("foobar"), remoteAddr, gomock.Any(), uint16(4), protocol.ECNCE).Return(0, errGSO), + rawConn.EXPECT().WritePacket([]byte("foob"), remoteAddr, gomock.Any(), uint16(0), protocol.ECNCE).Return(4, nil), + rawConn.EXPECT().WritePacket([]byte("ar"), remoteAddr, gomock.Any(), uint16(0), protocol.ECNCE).Return(2, nil), + ) + require.NoError(t, c.Write([]byte("foobar"), 4, protocol.ECNCE)) + require.False(t, c.capabilities().GSO) +} + +func TestSendConnSendmsgFailures(t *testing.T) { + if runtime.GOOS != "linux" { + t.Skip("only Linux exhibits this bug, we don't need to work around it on other platforms") + } + + remoteAddr := &net.UDPAddr{IP: net.IPv4(192, 168, 100, 200), Port: 1337} + + t.Run("first call to sendmsg fails", func(t *testing.T) { + rawConn := NewMockRawConn(gomock.NewController(t)) + rawConn.EXPECT().LocalAddr() + rawConn.EXPECT().capabilities().AnyTimes() + c := newSendConn(rawConn, remoteAddr, packetInfo{}, utils.DefaultLogger) + gomock.InOrder( + rawConn.EXPECT().WritePacket([]byte("foobar"), remoteAddr, gomock.Any(), gomock.Any(), protocol.ECNCE).Return(0, errNotPermitted), + rawConn.EXPECT().WritePacket([]byte("foobar"), remoteAddr, gomock.Any(), uint16(0), protocol.ECNCE).Return(6, nil), + ) + require.NoError(t, c.Write([]byte("foobar"), 0, protocol.ECNCE)) + }) + + t.Run("later call to sendmsg fails", func(t *testing.T) { + rawConn := NewMockRawConn(gomock.NewController(t)) + rawConn.EXPECT().LocalAddr() + rawConn.EXPECT().capabilities().AnyTimes() + c := newSendConn(rawConn, remoteAddr, packetInfo{}, utils.DefaultLogger) + rawConn.EXPECT().WritePacket([]byte("foobar"), remoteAddr, gomock.Any(), gomock.Any(), protocol.ECNCE).Return(0, errNotPermitted).Times(2) + require.Error(t, c.Write([]byte("foobar"), 0, protocol.ECNCE)) + }) +} + +func TestSendConnRemoteAddrChange(t *testing.T) { + ln1 := newUDPConnLocalhost(t) + ln2 := newUDPConnLocalhost(t) + + c := newSendConn( + &basicConn{PacketConn: newUDPConnLocalhost(t)}, + ln1.LocalAddr(), + packetInfo{}, + utils.DefaultLogger, + ) + + require.NoError(t, c.Write([]byte("foobar"), 0, protocol.ECNUnsupported)) + ln1.SetReadDeadline(time.Now().Add(time.Second)) + b := make([]byte, 1024) + n, err := ln1.Read(b) + require.NoError(t, err) + require.Equal(t, "foobar", string(b[:n])) + + require.NoError(t, c.WriteTo([]byte("foobaz"), ln2.LocalAddr(), packetInfo{})) + ln2.SetReadDeadline(time.Now().Add(time.Second)) + b = make([]byte, 1024) + n, err = ln2.Read(b) + require.NoError(t, err) + require.Equal(t, "foobaz", string(b[:n])) + + c.ChangeRemoteAddr(ln2.LocalAddr(), packetInfo{}) + require.NoError(t, c.Write([]byte("lorem ipsum"), 0, protocol.ECNUnsupported)) + ln2.SetReadDeadline(time.Now().Add(time.Second)) + b = make([]byte, 1024) + n, err = ln2.Read(b) + require.NoError(t, err) + require.Equal(t, "lorem ipsum", string(b[:n])) +} diff --git a/third_party/quic-go/send_queue.go b/third_party/quic-go/send_queue.go new file mode 100644 index 0000000..cabc1c4 --- /dev/null +++ b/third_party/quic-go/send_queue.go @@ -0,0 +1,114 @@ +package quic + +import ( + "errors" + "net" + + "github.com/apernet/quic-go/internal/protocol" +) + +type sender interface { + Send(p *packetBuffer, gsoSize uint16, ecn protocol.ECN) + SendProbe(*packetBuffer, net.Addr, packetInfo) + Run() error + WouldBlock() bool + Available() <-chan struct{} + Close() +} + +type queueEntry struct { + buf *packetBuffer + gsoSize uint16 + ecn protocol.ECN +} + +type sendQueue struct { + queue chan queueEntry + closeCalled chan struct{} // runStopped when Close() is called + runStopped chan struct{} // runStopped when the run loop returns + available chan struct{} + conn sendConn +} + +var _ sender = &sendQueue{} + +const sendQueueCapacity = 8 + +func newSendQueue(conn sendConn) sender { + return &sendQueue{ + conn: conn, + runStopped: make(chan struct{}), + closeCalled: make(chan struct{}), + available: make(chan struct{}, 1), + queue: make(chan queueEntry, sendQueueCapacity), + } +} + +// Send sends out a packet. It's guaranteed to not block. +// Callers need to make sure that there's actually space in the send queue by calling WouldBlock. +// Otherwise Send will panic. +func (h *sendQueue) Send(p *packetBuffer, gsoSize uint16, ecn protocol.ECN) { + select { + case h.queue <- queueEntry{buf: p, gsoSize: gsoSize, ecn: ecn}: + // clear available channel if we've reached capacity + if len(h.queue) == sendQueueCapacity { + select { + case <-h.available: + default: + } + } + case <-h.runStopped: + default: + panic("sendQueue.Send would have blocked") + } +} + +func (h *sendQueue) SendProbe(p *packetBuffer, addr net.Addr, info packetInfo) { + h.conn.WriteTo(p.Data, addr, info) +} + +func (h *sendQueue) WouldBlock() bool { + return len(h.queue) == sendQueueCapacity +} + +func (h *sendQueue) Available() <-chan struct{} { + return h.available +} + +func (h *sendQueue) Run() error { + defer close(h.runStopped) + var shouldClose bool + for { + if shouldClose && len(h.queue) == 0 { + return nil + } + select { + case <-h.closeCalled: + h.closeCalled = nil // prevent this case from being selected again + // make sure that all queued packets are actually sent out + shouldClose = true + case e := <-h.queue: + if err := h.conn.Write(e.buf.Data, e.gsoSize, e.ecn); err != nil { + // This additional check enables: + // 1. Checking for "datagram too large" message from the kernel, as such, + // 2. Path MTU discovery,and + // 3. Eventual detection of loss PingFrame. + var tooLarge *DatagramTooLargeError + if !isSendMsgSizeErr(err) && !errors.As(err, &tooLarge) { + return err + } + } + e.buf.Release() + select { + case h.available <- struct{}{}: + default: + } + } + } +} + +func (h *sendQueue) Close() { + close(h.closeCalled) + // wait until the run loop returned + <-h.runStopped +} diff --git a/third_party/quic-go/send_queue_test.go b/third_party/quic-go/send_queue_test.go new file mode 100644 index 0000000..1e05bcd --- /dev/null +++ b/third_party/quic-go/send_queue_test.go @@ -0,0 +1,207 @@ +package quic + +import ( + "net" + "net/netip" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func getPacketWithContents(b []byte) *packetBuffer { + buf := getPacketBuffer() + buf.Data = buf.Data[:len(b)] + copy(buf.Data, b) + return buf +} + +func TestSendQueueSendOnePacket(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + c := NewMockSendConn(mockCtrl) + q := newSendQueue(c) + + written := make(chan struct{}) + c.EXPECT().Write([]byte("foobar"), uint16(10), protocol.ECT1).Do( + func([]byte, uint16, protocol.ECN) error { close(written); return nil }, + ) + + done := make(chan struct{}) + go func() { + q.Run() + close(done) + }() + + q.Send(getPacketWithContents([]byte("foobar")), 10, protocol.ECT1) + synctest.Wait() + + select { + case <-written: + default: + t.Fatal("write should have returned") + } + + q.Close() + synctest.Wait() + + select { + case <-done: + default: + t.Fatal("Run should have returned") + } + }) +} + +func TestSendQueueBlocking(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + c := NewMockSendConn(mockCtrl) + q := newSendQueue(c) + + blockWrite := make(chan struct{}) + written := make(chan struct{}, 1) + c.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).Do( + func([]byte, uint16, protocol.ECN) error { + select { + case written <- struct{}{}: + default: + } + <-blockWrite + return nil + }, + ).AnyTimes() + + done := make(chan struct{}) + go func() { + q.Run() + close(done) + }() + + // +1, since one packet will be queued in the Write call + for i := range sendQueueCapacity + 1 { + require.False(t, q.WouldBlock()) + q.Send(getPacketWithContents([]byte("foobar")), 10, protocol.ECT1) + // make sure that the first packet is actually enqueued in the Write call + if i == 0 { + select { + case <-written: + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + } + require.True(t, q.WouldBlock()) + select { + case <-q.Available(): + t.Fatal("should not be available") + default: + } + require.Panics(t, func() { q.Send(getPacketWithContents([]byte("foobar")), 10, protocol.ECT1) }) + + // allow one packet to be sent + blockWrite <- struct{}{} + select { + case <-written: + case <-time.After(time.Second): + t.Fatal("timeout") + } + select { + case <-q.Available(): + require.False(t, q.WouldBlock()) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // when calling Close, all packets are first sent out + closed := make(chan struct{}) + go func() { + q.Close() + close(closed) + }() + + synctest.Wait() + + select { + case <-closed: + t.Fatal("Close should have blocked") + default: + } + + for range sendQueueCapacity { + blockWrite <- struct{}{} + } + synctest.Wait() + + select { + case <-closed: + default: + t.Fatal("Close should have returned") + } + select { + case <-done: + default: + t.Fatal("Run should have returned") + } + }) +} + +func TestSendQueueWriteError(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + c := NewMockSendConn(mockCtrl) + q := newSendQueue(c) + + c.EXPECT().Write(gomock.Any(), gomock.Any(), gomock.Any()).Return(assert.AnError) + q.Send(getPacketWithContents([]byte("foobar")), 6, protocol.ECNNon) + + errChan := make(chan error, 1) + go func() { errChan <- q.Run() }() + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, assert.AnError) + default: + t.Fatal("Run should have returned") + } + + // further calls to Send should not block + sent := make(chan struct{}) + go func() { + defer close(sent) + for range 2 * sendQueueCapacity { + q.Send(getPacketWithContents([]byte("raboof")), 6, protocol.ECNNon) + } + }() + + synctest.Wait() + + select { + case <-sent: + default: + t.Fatal("Send should have returned") + } + }) +} + +func TestSendQueueSendProbe(t *testing.T) { + mockCtrl := gomock.NewController(t) + c := NewMockSendConn(mockCtrl) + q := newSendQueue(c) + + addr := &net.UDPAddr{IP: net.IPv4(42, 42, 42, 42), Port: 42} + localAddr := netip.MustParseAddr("43.43.43.43") + c.EXPECT().WriteTo([]byte("foobar"), addr, packetInfo{ + addr: localAddr, + }) + q.SendProbe(getPacketWithContents([]byte("foobar")), addr, packetInfo{ + addr: localAddr, + }) +} diff --git a/third_party/quic-go/send_stream.go b/third_party/quic-go/send_stream.go new file mode 100644 index 0000000..65cffda --- /dev/null +++ b/third_party/quic-go/send_stream.go @@ -0,0 +1,915 @@ +package quic + +import ( + "context" + "fmt" + "sync" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +// A SendStream is a unidirectional Send Stream. +type SendStream struct { + mutex sync.Mutex + + numOutstandingFrames int64 // outstanding STREAM and RESET_STREAM frames + retransmissionQueue []*wire.StreamFrame + + ctx context.Context + ctxCancel context.CancelCauseFunc + + streamID protocol.StreamID + sender streamSender + + // reliableSize is the portion of the stream that needs to be transmitted reliably, + // even if the stream is cancelled. + // This requires the peer to support RESET_STREAM_AT. + // This value should not be accessed directly, but only through the reliableOffset method. + // This method returns 0 if the peer doesn't support the RESET_STREAM_AT extension. + reliableSize protocol.ByteCount + writeOffset protocol.ByteCount + + shutdownErr error + resetErr *StreamError + queuedResetStreamFrame *wire.ResetStreamFrame + + dataForWriting []byte // during a Write() call, this slice is the part of p that still needs to be sent out + + writeLimiter func(int) int + // Set by the packetizer when writeLimiter reduces the allowed byte count. It makes the blocked + // WriteWithLimit return ErrWriteLimitReached and prevents another dequeue before it wakes up. + writeLimited bool + + nextFrame *wire.StreamFrame + // set if flow control credit for nextFrame was already consumed + nextFrameReserved bool + + supportsResetStreamAt bool + finishedWriting bool // set once Close() is called + finSent bool // set when a STREAM_FRAME with FIN bit has been sent + // Set when the application knows about the cancellation. + // This can happen because the application called CancelWrite, + // or because Write returned the error (for remote cancellations). + cancellationFlagged bool + completed bool // set when this stream has been reported to the streamSender as completed + + writeChan chan struct{} + writeOnce chan struct{} + deadline monotime.Time + + flowController *streamFlowController +} + +var ( + _ streamControlFrameGetter = &SendStream{} + _ outgoingStream = &SendStream{} + _ sendStreamFrameHandler = &SendStream{} +) + +func newSendStream( + ctx context.Context, + streamID protocol.StreamID, + sender streamSender, + flowController *streamFlowController, + supportsResetStreamAt bool, +) *SendStream { + s := &SendStream{ + streamID: streamID, + sender: sender, + flowController: flowController, + writeChan: make(chan struct{}, 1), + writeOnce: make(chan struct{}, 1), // cap: 1, to protect against concurrent use of Write + supportsResetStreamAt: supportsResetStreamAt, + } + s.ctx, s.ctxCancel = context.WithCancelCause(ctx) + return s +} + +// StreamID returns the stream ID. +func (s *SendStream) StreamID() StreamID { + return s.streamID // same for receiveStream and sendStream +} + +// Write writes data to the stream. +// Write can be made to time out using [SendStream.SetWriteDeadline]. +// If the stream was canceled, the error is a [StreamError]. +func (s *SendStream) Write(p []byte) (int, error) { + return s.WriteWithLimit(p, nil) +} + +// WriteWithLimit writes data to the stream, subject to an additional send limit. +// During packetization, limiter receives the bytes allowed for the next STREAM frame after +// QUIC flow control and returns how many may be sent. Returning n in [0, maxBytes] commits +// n bytes of limiter credit; the limiter is not called again when those bytes are retransmitted. +// Values outside [0, maxBytes] are clamped. +// A short result returns the accepted prefix and [ErrWriteLimitReached]; the caller can wait +// for external credit and retry the suffix. QUIC blocking behaves like [SendStream.Write]. +// limiter can run multiple times on another goroutine while QUIC send flow-control accounting +// is locked. It must be concurrency-safe and must not block or call QUIC methods. +// A nil limiter behaves like [SendStream.Write]. +func (s *SendStream) WriteWithLimit(p []byte, limiter func(maxBytes int) int) (int, error) { + // Concurrent use of Write is not permitted (and doesn't make any sense), + // but sometimes people do it anyway. + // Make sure that we only execute one call at any given time to avoid hard to debug failures. + s.writeOnce <- struct{}{} + defer func() { <-s.writeOnce }() + + isNewlyCompleted, n, err := s.write(p, limiter) + if isNewlyCompleted { + s.sender.onStreamCompleted(s.streamID) + } + return n, err +} + +// TryWriteAll writes data to the stream if it can be queued immediately. +// It doesn't block for flow control credit and doesn't respect the write deadline. +// If the entire slice can't be queued immediately, it queues nothing and returns [ErrWouldBlock]. +func (s *SendStream) TryWriteAll(p []byte) error { + select { + case s.writeOnce <- struct{}{}: + defer func() { <-s.writeOnce }() + default: + return ErrWouldBlock + } + + isNewlyCompleted, hasData, err := s.tryWriteAll(p) + if isNewlyCompleted { + s.sender.onStreamCompleted(s.streamID) + } + if hasData { + s.sender.onHasStreamData(s.streamID, s) + } + return err +} + +func (s *SendStream) tryWriteAll(p []byte) (bool /* is newly completed */, bool /* has data */, error) { + // This might wait briefly while a packet is dequeuing stream data. + s.mutex.Lock() + defer s.mutex.Unlock() + + if s.resetErr != nil { + s.cancellationFlagged = true + return s.isNewlyCompleted(), false, s.resetErr + } + if s.shutdownErr != nil { + return false, false, s.shutdownErr + } + if s.finishedWriting { + return false, false, fmt.Errorf("write on closed stream %d", s.streamID) + } + if len(p) == 0 { + return false, false, nil + } + + bytesToReserve := protocol.ByteCount(len(p)) + if s.nextFrame != nil && !s.nextFrameReserved { + bytesToReserve += s.nextFrame.DataLen() + } + if !s.flowController.TryAddBytesSent(bytesToReserve) { + return false, false, ErrWouldBlock + } + + if s.nextFrame == nil { + s.nextFrame = wire.GetStreamFrame() + s.nextFrame.Offset = s.writeOffset + s.nextFrame.StreamID = s.streamID + s.nextFrame.DataLenPresent = true + s.nextFrame.Data = s.nextFrame.Data[:0] + } + l := len(s.nextFrame.Data) + if l+len(p) > cap(s.nextFrame.Data) { + // Pooled STREAM frames must keep their packet-sized buffer. + // Use a non-pooled frame when the queued data grows beyond that. + nextFrame := &wire.StreamFrame{ + StreamID: s.streamID, + Offset: s.nextFrame.Offset, + DataLenPresent: true, + Data: make([]byte, l+len(p)), + } + copy(nextFrame.Data, s.nextFrame.Data) + s.nextFrame.PutBack() + s.nextFrame = nextFrame + } else { + s.nextFrame.Data = s.nextFrame.Data[:l+len(p)] + } + copy(s.nextFrame.Data[l:], p) + s.nextFrameReserved = true + return false, true, nil +} + +func (s *SendStream) write(p []byte, limiter func(int) int) (bool /* is newly completed */, int, error) { + s.mutex.Lock() + s.writeLimiter = limiter + s.writeLimited = false + defer func() { + s.writeLimiter = nil + s.writeLimited = false + s.mutex.Unlock() + }() + + if s.resetErr != nil { + s.cancellationFlagged = true + return s.isNewlyCompleted(), 0, s.resetErr + } + if s.shutdownErr != nil { + return false, 0, s.shutdownErr + } + if s.finishedWriting { + return false, 0, fmt.Errorf("write on closed stream %d", s.streamID) + } + if !s.deadline.IsZero() && !monotime.Now().Before(s.deadline) { + return false, 0, errDeadline + } + if len(p) == 0 { + return false, 0, nil + } + + s.dataForWriting = p + + var ( + deadlineTimer *time.Timer + bytesWritten int + notifiedSender bool + ) + for { + if s.writeLimited { + bytesWritten = len(p) - len(s.dataForWriting) + s.dataForWriting = nil + break + } + var copied bool + var deadline monotime.Time + // As soon as dataForWriting becomes smaller than a certain size x, we copy all the data to a STREAM frame (s.nextFrame), + // which can then be popped the next time we assemble a packet. + // This allows us to return Write() when all data but x bytes have been sent out. + // When the user now calls Close(), this is much more likely to happen before we popped that last STREAM frame, + // allowing us to set the FIN bit on that frame (instead of sending an empty STREAM frame with FIN). + if s.canBufferStreamFrame() && len(s.dataForWriting) > 0 { + if s.nextFrame == nil { + f := wire.GetStreamFrame() + f.Offset = s.writeOffset + f.StreamID = s.streamID + f.DataLenPresent = true + f.Data = f.Data[:len(s.dataForWriting)] + copy(f.Data, s.dataForWriting) + s.nextFrame = f + } else { + l := len(s.nextFrame.Data) + s.nextFrame.Data = s.nextFrame.Data[:l+len(s.dataForWriting)] + copy(s.nextFrame.Data[l:], s.dataForWriting) + } + s.dataForWriting = nil + bytesWritten = len(p) + copied = true + } else { + bytesWritten = len(p) - len(s.dataForWriting) + deadline = s.deadline + if !deadline.IsZero() { + if !monotime.Now().Before(deadline) { + s.dataForWriting = nil + return false, bytesWritten, errDeadline + } + if deadlineTimer == nil { + deadlineTimer = time.NewTimer(monotime.Until(deadline)) + defer deadlineTimer.Stop() + } else { + deadlineTimer.Reset(monotime.Until(deadline)) + } + } + if s.dataForWriting == nil || s.shutdownErr != nil || s.resetErr != nil { + break + } + } + + s.mutex.Unlock() + if !notifiedSender { + s.sender.onHasStreamData(s.streamID, s) // must be called without holding the mutex + notifiedSender = true + } + if copied { + s.mutex.Lock() + break + } + if deadline.IsZero() { + <-s.writeChan + } else { + select { + case <-s.writeChan: + case <-deadlineTimer.C: + } + } + s.mutex.Lock() + } + + if bytesWritten == len(p) { + return false, bytesWritten, nil + } + if s.shutdownErr != nil { + return false, bytesWritten, s.shutdownErr + } + if s.resetErr != nil { + s.cancellationFlagged = true + return s.isNewlyCompleted(), bytesWritten, s.resetErr + } + if s.writeLimited { + return false, bytesWritten, ErrWriteLimitReached + } + return false, bytesWritten, nil +} + +func (s *SendStream) canBufferStreamFrame() bool { + if s.writeLimiter != nil || s.nextFrameReserved { + return false + } + var l protocol.ByteCount + if s.nextFrame != nil { + l = s.nextFrame.DataLen() + } + return l+protocol.ByteCount(len(s.dataForWriting)) <= protocol.MaxPacketBufferSize +} + +// popStreamFrame returns the next STREAM frame that is supposed to be sent on this stream +// maxBytes is the maximum length this frame (including frame header) will have. +func (s *SendStream) popStreamFrame(maxBytes protocol.ByteCount, v protocol.Version) (_ ackhandler.StreamFrame, _ *wire.StreamDataBlockedFrame, hasMore bool) { + s.mutex.Lock() + f, blocked, hasMoreData := s.popNewOrRetransmittedStreamFrame(maxBytes, v) + if f != nil { + s.numOutstandingFrames++ + } + s.mutex.Unlock() + + if f == nil { + return ackhandler.StreamFrame{}, blocked, hasMoreData + } + return ackhandler.StreamFrame{ + Frame: f, + Handler: (*sendStreamAckHandler)(s), + }, blocked, hasMoreData +} + +func (s *SendStream) popNewOrRetransmittedStreamFrame(maxBytes protocol.ByteCount, v protocol.Version) (_ *wire.StreamFrame, _ *wire.StreamDataBlockedFrame, hasMoreData bool) { + if s.shutdownErr != nil { + return nil, nil, false + } + if s.resetErr != nil { + reliableOffset := s.reliableOffset() + if reliableOffset == 0 || (s.writeOffset >= reliableOffset && len(s.retransmissionQueue) == 0) { + return nil, nil, false + } + } + + if len(s.retransmissionQueue) > 0 { + f, hasMoreRetransmissions := s.maybeGetRetransmission(maxBytes, v) + if f != nil || hasMoreRetransmissions { + if f == nil { + return nil, nil, true + } + // We always claim that we have more data to send. + // This might be incorrect, in which case there'll be a spurious call to popStreamFrame in the future. + return f, nil, true + } + } + if s.writeLimited { + return nil, nil, false + } + + if len(s.dataForWriting) == 0 && s.nextFrame == nil { + if s.finishedWriting && !s.finSent { + s.finSent = true + return &wire.StreamFrame{ + StreamID: s.streamID, + Offset: s.writeOffset, + DataLenPresent: true, + Fin: true, + }, nil, false + } + return nil, nil, false + } + + // if the stream is canceled, only data up to the reliable size needs to be sent + reliableOffset := s.reliableOffset() + limitedWrite := s.writeLimiter != nil && s.nextFrame == nil + var maxDataLen protocol.ByteCount + if s.nextFrameReserved { + maxDataLen = s.nextFrame.DataLen() + } else { + maxDataLen = s.flowController.SendWindowSize() + } + if s.resetErr != nil && reliableOffset > 0 { + maxDataLen = min(maxDataLen, reliableOffset-s.writeOffset) + } + if s.nextFrame != nil { + maxDataLen = min(maxDataLen, s.nextFrame.MaxDataLen(maxBytes, v), s.nextFrame.DataLen()) + } else { + f := wire.StreamFrame{ + StreamID: s.streamID, + Offset: s.writeOffset, + DataLenPresent: true, + } + maxDataLen = min(maxDataLen, f.MaxDataLen(maxBytes, v), protocol.ByteCount(len(s.dataForWriting))) + } + if maxDataLen == 0 { + return nil, nil, true + } + if limitedWrite { + added, limited := s.flowController.AddBytesSentWithLimiter(maxDataLen, s.writeLimiter) + if limited { + s.writeLimited = true + s.signalWrite() + } + maxDataLen = added + if maxDataLen == 0 { + return nil, nil, !limited + } + } else if !s.nextFrameReserved && !s.flowController.TryAddBytesSent(maxDataLen) { + return nil, nil, true + } + f, hasMoreData := s.popNewStreamFrame(maxDataLen) + if f.DataLen() > 0 { + s.writeOffset += f.DataLen() + } + if s.resetErr != nil && s.writeOffset >= reliableOffset { + hasMoreData = false + } + if s.writeLimited { + hasMoreData = false + } + var blocked *wire.StreamDataBlockedFrame + if f.DataLen() > 0 { + if isBlocked, offset := s.flowController.isNewlyBlocked(); isBlocked { + blocked = &wire.StreamDataBlockedFrame{StreamID: s.streamID, MaximumStreamData: offset} + } + } + f.Fin = s.finishedWriting && s.dataForWriting == nil && s.nextFrame == nil && !s.finSent + if f.Fin { + s.finSent = true + } + return f, blocked, hasMoreData +} + +// popNewStreamFrame returns a new STREAM frame to send for this stream +// hasMoreData says if there's more data to send, *not* taking into account the reliable size +func (s *SendStream) popNewStreamFrame(maxDataLen protocol.ByteCount) (_ *wire.StreamFrame, hasMoreData bool) { + if s.nextFrame != nil { + nextFrame := s.nextFrame + nextFrameReserved := s.nextFrameReserved + s.nextFrame = nil + s.nextFrameReserved = false + if nextFrame.DataLen() > maxDataLen { + if nextFrame.DataLen()-maxDataLen > protocol.MaxPacketBufferSize { + s.nextFrame = &wire.StreamFrame{ + Data: make([]byte, nextFrame.DataLen()-maxDataLen), + } + } else { + s.nextFrame = wire.GetStreamFrame() + s.nextFrame.Data = s.nextFrame.Data[:nextFrame.DataLen()-maxDataLen] + } + s.nextFrame.StreamID = s.streamID + s.nextFrame.Offset = s.writeOffset + maxDataLen + s.nextFrame.DataLenPresent = true + copy(s.nextFrame.Data, nextFrame.Data[maxDataLen:]) + nextFrame.Data = nextFrame.Data[:maxDataLen] + s.nextFrameReserved = nextFrameReserved + } else { + s.signalWrite() + } + return nextFrame, s.nextFrame != nil || s.dataForWriting != nil + } + + f := wire.GetStreamFrame() + f.Fin = false + f.StreamID = s.streamID + f.Offset = s.writeOffset + f.DataLenPresent = true + f.Data = f.Data[:0] + + s.getDataForWriting(f, maxDataLen) + return f, s.dataForWriting != nil || s.nextFrame != nil || s.finishedWriting +} + +func (s *SendStream) maybeGetRetransmission(maxBytes protocol.ByteCount, v protocol.Version) (*wire.StreamFrame, bool /* has more retransmissions */) { + f := s.retransmissionQueue[0] + newFrame, needsSplit := f.MaybeSplitOffFrame(maxBytes, v) + if needsSplit { + return newFrame, true + } + s.retransmissionQueue = s.retransmissionQueue[1:] + return f, len(s.retransmissionQueue) > 0 +} + +func (s *SendStream) getDataForWriting(f *wire.StreamFrame, maxBytes protocol.ByteCount) { + if protocol.ByteCount(len(s.dataForWriting)) <= maxBytes { + f.Data = f.Data[:len(s.dataForWriting)] + copy(f.Data, s.dataForWriting) + s.dataForWriting = nil + s.signalWrite() + return + } + f.Data = f.Data[:maxBytes] + copy(f.Data, s.dataForWriting) + s.dataForWriting = s.dataForWriting[maxBytes:] + if s.canBufferStreamFrame() { + s.signalWrite() + } +} + +func (s *SendStream) isNewlyCompleted() bool { + if s.completed { + return false + } + if s.nextFrame != nil && s.nextFrame.DataLen() > 0 { + return false + } + // We need to keep the stream around until all frames have been sent and acknowledged. + if s.numOutstandingFrames > 0 || len(s.retransmissionQueue) > 0 || s.queuedResetStreamFrame != nil { + return false + } + // The stream is completed if we sent the FIN. + if s.finSent { + s.completed = true + return true + } + // The stream is also completed if: + // 1. the application called CancelWrite, or + // 2. we received a STOP_SENDING, and + // * the application consumed the error via Write, or + // * the application called Close + if s.resetErr != nil && (s.cancellationFlagged || s.finishedWriting) { + s.completed = true + return true + } + return false +} + +// Close closes the write-direction of the stream. +// Future calls to Write are not permitted after calling Close. +// It must not be called concurrently with Write. +// It must not be called after calling CancelWrite. +func (s *SendStream) Close() error { + s.mutex.Lock() + if s.shutdownErr != nil || s.finishedWriting { + s.mutex.Unlock() + return nil + } + s.finishedWriting = true + cancelled := s.resetErr != nil + if cancelled { + s.cancellationFlagged = true + } + completed := s.isNewlyCompleted() + s.mutex.Unlock() + + if completed { + s.sender.onStreamCompleted(s.streamID) + } + if cancelled { + return fmt.Errorf("close called for canceled stream %d", s.streamID) + } + s.sender.onHasStreamData(s.streamID, s) // need to send the FIN, must be called without holding the mutex + + s.ctxCancel(nil) + return nil +} + +// SetReliableBoundary marks the data written to this stream so far as reliable. +// It is valid to call this function multiple times, thereby increasing the reliable size. +// It only has an effect if the peer enabled support for the RESET_STREAM_AT extension, +// otherwise, it is a no-op. +func (s *SendStream) SetReliableBoundary() { + s.mutex.Lock() + defer s.mutex.Unlock() + + if s.nextFrame != nil { + s.reliableSize = max(s.reliableSize, s.writeOffset+s.nextFrame.DataLen()) + } else { + s.reliableSize = max(s.reliableSize, s.writeOffset) + } +} + +// returnFramesToPool returns all queued frames to the sync.Pool +func (s *SendStream) returnFramesToPool() { + for _, f := range s.retransmissionQueue { + f.PutBack() + } + clear(s.retransmissionQueue) + s.retransmissionQueue = nil + if s.nextFrame != nil { + s.nextFrame.PutBack() + s.nextFrame = nil + } + s.nextFrameReserved = false +} + +// CancelWrite aborts sending on this stream. +// Data already written, but not yet delivered to the peer is not guaranteed to be delivered reliably. +// Write will unblock immediately, and future calls to Write will fail. +// When called multiple times it is a no-op. +// When called after Close, it aborts reliable delivery of outstanding stream data. +// Note that there is no guarantee if the peer will receive the FIN or the cancellation error first. +func (s *SendStream) CancelWrite(errorCode StreamErrorCode) { + s.mutex.Lock() + if s.shutdownErr != nil { + s.mutex.Unlock() + return + } + + s.cancellationFlagged = true + + if s.resetErr != nil { + completed := s.isNewlyCompleted() + s.mutex.Unlock() + // The user has called CancelWrite. If the previous cancellation was because of a + // STOP_SENDING, we don't need to flag the error to the user anymore. + if completed { + s.sender.onStreamCompleted(s.streamID) + } + return + } + s.resetErr = &StreamError{StreamID: s.streamID, ErrorCode: errorCode, Remote: false} + s.ctxCancel(s.resetErr) + + reliableOffset := s.reliableOffset() + finalSize := max(s.writeOffset, reliableOffset) + if s.nextFrameReserved && s.nextFrame != nil { + finalSize = max(finalSize, s.nextFrame.Offset+s.nextFrame.DataLen()) + } + if reliableOffset == 0 { + s.numOutstandingFrames = 0 + s.returnFramesToPool() + } + s.queuedResetStreamFrame = &wire.ResetStreamFrame{ + StreamID: s.streamID, + FinalSize: finalSize, + ErrorCode: errorCode, + // if the peer doesn't support the extension, the reliable offset will always be 0 + ReliableSize: reliableOffset, + } + if reliableOffset > 0 { + if s.nextFrame != nil { + if s.nextFrame.Offset >= reliableOffset { + s.nextFrame.PutBack() + s.nextFrame = nil + s.nextFrameReserved = false + } else if s.nextFrame.Offset+s.nextFrame.DataLen() > reliableOffset { + s.nextFrame.Data = s.nextFrame.Data[:reliableOffset-s.nextFrame.Offset] + } + } + if len(s.retransmissionQueue) > 0 { + retransmissionQueue := make([]*wire.StreamFrame, 0, len(s.retransmissionQueue)) + for _, f := range s.retransmissionQueue { + if f.Offset >= reliableOffset { + f.PutBack() + continue + } + if f.Offset+f.DataLen() <= reliableOffset { + retransmissionQueue = append(retransmissionQueue, f) + } else { + f.Data = f.Data[:reliableOffset-f.Offset] + retransmissionQueue = append(retransmissionQueue, f) + } + } + s.retransmissionQueue = retransmissionQueue + } + } + s.mutex.Unlock() + + s.signalWrite() + s.sender.onHasStreamControlFrame(s.streamID, s) +} + +func (s *SendStream) enableResetStreamAt() { + s.mutex.Lock() + s.supportsResetStreamAt = true + s.mutex.Unlock() +} + +func (s *SendStream) updateSendWindow(limit protocol.ByteCount) { + s.mutex.Lock() + updated := s.flowController.UpdateSendWindow(limit) + if !updated { // duplicate or reordered MAX_STREAM_DATA frame + s.mutex.Unlock() + return + } + hasStreamData := s.dataForWriting != nil || s.nextFrame != nil + s.mutex.Unlock() + if hasStreamData { + s.sender.onHasStreamData(s.streamID, s) + } +} + +func (s *SendStream) handleStopSendingFrame(f *wire.StopSendingFrame) { + s.mutex.Lock() + if s.shutdownErr != nil { + s.mutex.Unlock() + return + } + + // If the stream was already cancelled (either locally, or due to a previous STOP_SENDING frame), + // there's nothing else to do. + if s.resetErr != nil && s.reliableOffset() == 0 { + s.mutex.Unlock() + return + } + // if the peer stopped reading from the stream, there's no need to transmit any data reliably + s.reliableSize = 0 + s.numOutstandingFrames = 0 + finalSize := s.writeOffset + if s.nextFrameReserved && s.nextFrame != nil { + finalSize = max(finalSize, s.nextFrame.Offset+s.nextFrame.DataLen()) + } + s.returnFramesToPool() + if s.resetErr == nil { + s.resetErr = &StreamError{StreamID: s.streamID, ErrorCode: f.ErrorCode, Remote: true} + s.ctxCancel(s.resetErr) + } + s.queuedResetStreamFrame = &wire.ResetStreamFrame{ + StreamID: s.streamID, + FinalSize: finalSize, + ErrorCode: s.resetErr.ErrorCode, + } + s.mutex.Unlock() + + s.signalWrite() + s.sender.onHasStreamControlFrame(s.streamID, s) +} + +func (s *SendStream) getControlFrame(monotime.Time) (_ ackhandler.Frame, ok, hasMore bool) { + s.mutex.Lock() + defer s.mutex.Unlock() + + if s.queuedResetStreamFrame == nil { + return ackhandler.Frame{}, false, false + } + s.numOutstandingFrames++ + f := ackhandler.Frame{ + Frame: s.queuedResetStreamFrame, + Handler: (*sendStreamResetStreamHandler)(s), + } + s.queuedResetStreamFrame = nil + return f, true, false +} + +func (s *SendStream) reliableOffset() protocol.ByteCount { + if !s.supportsResetStreamAt { + return 0 + } + return s.reliableSize +} + +// The Context is canceled as soon as the write-side of the stream is closed. +// This happens when Close() or CancelWrite() is called, or when the peer +// cancels the read-side of their stream. +// The cancellation cause is set to the error that caused the stream to +// close, or `context.Canceled` in case the stream is closed without error. +func (s *SendStream) Context() context.Context { + return s.ctx +} + +// SetWriteDeadline sets the deadline for future Write calls +// and any currently-blocked Write call. +// Even if write times out, it may return n > 0, indicating that +// some data was successfully written. +// A zero value for t means Write will not time out. +func (s *SendStream) SetWriteDeadline(t time.Time) error { + s.mutex.Lock() + s.deadline = monotime.FromTime(t) + s.mutex.Unlock() + s.signalWrite() + return nil +} + +// CloseForShutdown closes a stream abruptly. +// It makes Write unblock (and return the error) immediately. +// The peer will NOT be informed about this: the stream is closed without sending a FIN or RST. +func (s *SendStream) closeForShutdown(err error) { + s.mutex.Lock() + if s.shutdownErr == nil && !s.finishedWriting { + s.shutdownErr = err + s.returnFramesToPool() + } + s.mutex.Unlock() + s.ctxCancel(err) + s.signalWrite() +} + +// signalWrite performs a non-blocking send on the writeChan +func (s *SendStream) signalWrite() { + select { + case s.writeChan <- struct{}{}: + default: + } +} + +type sendStreamAckHandler SendStream + +var _ ackhandler.FrameHandler = &sendStreamAckHandler{} + +func (s *sendStreamAckHandler) OnAcked(f wire.Frame) { + sf := f.(*wire.StreamFrame) + sf.PutBack() + + s.mutex.Lock() + if s.resetErr != nil && (*SendStream)(s).reliableOffset() == 0 { + s.mutex.Unlock() + return + } + s.numOutstandingFrames-- + if s.numOutstandingFrames < 0 { + panic("numOutStandingFrames negative") + } + completed := (*SendStream)(s).isNewlyCompleted() + s.mutex.Unlock() + + if completed { + s.sender.onStreamCompleted(s.streamID) + } +} + +func (s *sendStreamAckHandler) OnLost(f wire.Frame) { + sf := f.(*wire.StreamFrame) + s.mutex.Lock() + // If the reliable size was 0 when the stream was cancelled, + // the number of outstanding frames was immediately set to 0, and the retransmission queue was dropped. + if s.resetErr != nil && (*SendStream)(s).reliableOffset() == 0 { + // Return the frame to pool since it won't be retransmitted + sf.PutBack() + s.mutex.Unlock() + return + } + s.numOutstandingFrames-- + if s.numOutstandingFrames < 0 { + panic("numOutStandingFrames negative") + } + + if s.resetErr != nil && (*SendStream)(s).reliableOffset() > 0 { + // If the stream was reset, and this frame is beyond the reliable offset, + // it doesn't need to be retransmitted. + if sf.Offset >= (*SendStream)(s).reliableOffset() { + sf.PutBack() + // If this frame was the last one tracked, losing it might cause the stream to be completed. + completed := (*SendStream)(s).isNewlyCompleted() + s.mutex.Unlock() + if completed { + s.sender.onStreamCompleted(s.streamID) + } + return + } + // If the payload of the frame extends beyond the reliable size, + // truncate the frame to the reliable size. + if sf.Offset+sf.DataLen() > (*SendStream)(s).reliableOffset() { + sf.Data = sf.Data[:(*SendStream)(s).reliableOffset()-sf.Offset] + } + } + + sf.DataLenPresent = true + s.retransmissionQueue = append(s.retransmissionQueue, sf) + s.mutex.Unlock() + + s.sender.onHasStreamData(s.streamID, (*SendStream)(s)) +} + +type sendStreamResetStreamHandler SendStream + +var _ ackhandler.FrameHandler = &sendStreamResetStreamHandler{} + +func (s *sendStreamResetStreamHandler) OnAcked(f wire.Frame) { + rsf := f.(*wire.ResetStreamFrame) + s.mutex.Lock() + // If the peer sent a STOP_SENDING after we sent a RESET_STREAM_AT frame, + // we sent 1. reduced the reliable size to 0 and 2. sent a RESET_STREAM frame. + // In this case, we don't care about the acknowledgment of this frame. + if rsf.ReliableSize != (*SendStream)(s).reliableOffset() { + s.mutex.Unlock() + return + } + s.numOutstandingFrames-- + if s.numOutstandingFrames < 0 { + panic("numOutStandingFrames negative") + } + completed := (*SendStream)(s).isNewlyCompleted() + s.mutex.Unlock() + + if completed { + s.sender.onStreamCompleted(s.streamID) + } +} + +func (s *sendStreamResetStreamHandler) OnLost(f wire.Frame) { + rsf := f.(*wire.ResetStreamFrame) + s.mutex.Lock() + // If the peer sent a STOP_SENDING after we sent a RESET_STREAM_AT frame, + // we sent 1. reduced the reliable size to 0 and 2. sent a RESET_STREAM frame. + // In this case, the loss of the RESET_STREAM_AT frame can be ignored. + if rsf.ReliableSize != (*SendStream)(s).reliableOffset() { + s.mutex.Unlock() + return + } + s.queuedResetStreamFrame = rsf + s.numOutstandingFrames-- + s.mutex.Unlock() + s.sender.onHasStreamControlFrame(s.streamID, (*SendStream)(s)) +} diff --git a/third_party/quic-go/send_stream_test.go b/third_party/quic-go/send_stream_test.go new file mode 100644 index 0000000..a602693 --- /dev/null +++ b/third_party/quic-go/send_stream_test.go @@ -0,0 +1,1875 @@ +package quic + +import ( + "bytes" + "context" + "crypto/rand" + "errors" + "fmt" + "io" + mrand "math/rand/v2" + "net" + "os" + "runtime" + "slices" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type writerWithTimeout struct { + io.Writer + Timeout time.Duration +} + +func (w *writerWithTimeout) Write(p []byte) (n int, err error) { + done := make(chan struct{}) + go func() { + defer close(done) + n, err = w.Writer.Write(p) + }() + + select { + case <-done: + return n, err + case <-time.After(w.Timeout): + return 0, fmt.Errorf("write timeout after %s", w.Timeout) + } +} + +func expectedFrameHeaderLen(strID protocol.StreamID, offset protocol.ByteCount) protocol.ByteCount { + return (&wire.StreamFrame{StreamID: strID, Offset: offset, DataLenPresent: true}).Length(protocol.Version1) +} + +func TestSendStreamSetup(t *testing.T) { + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + ctx := context.WithValue(context.Background(), "foo", "bar") + str := newSendStream(ctx, 1337, nil, mockFC, false) + require.NotNil(t, str.Context()) + require.Equal(t, "bar", str.Context().Value("foo")) + require.Equal(t, protocol.StreamID(1337), str.StreamID()) +} + +func TestSendStreamWriteData(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + strWithTimeout := &writerWithTimeout{Writer: str, Timeout: time.Second} + + mockSender.EXPECT().onHasStreamData(streamID, str) + n, err := strWithTimeout.Write([]byte("foobar")) + require.NoError(t, err) + require.Equal(t, 6, n) + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foobar"), DataLenPresent: true}, + frame.Frame, + ) + require.True(t, mockCtrl.Satisfied()) + + // nothing more to send at this point + _, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // nil writes don't do anything + n, err = strWithTimeout.Write(nil) + require.NoError(t, err) + require.Zero(t, n) + require.True(t, mockCtrl.Satisfied()) + + // empty slices writes don't do anything + n, err = strWithTimeout.Write([]byte{}) + require.NoError(t, err) + require.Zero(t, n) + require.True(t, mockCtrl.Satisfied()) + + // multiple writes are bundled into a single frame + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + n, err = strWithTimeout.Write([]byte{0xde, 0xad}) + require.NoError(t, err) + require.Equal(t, 2, n) + n, err = strWithTimeout.Write([]byte{0xbe, 0xef}) + require.NoError(t, err) + require.Equal(t, 2, n) + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 42, Offset: 6, Data: []byte{0xde, 0xad, 0xbe, 0xef}, DataLenPresent: true}, + frame.Frame, + ) + + // a single write is split up into smaller frames + mockSender.EXPECT().onHasStreamData(streamID, str) + n, err = strWithTimeout.Write([]byte("foobaz")) + require.NoError(t, err) + require.Equal(t, 6, n) + frame, _, hasMore = str.popStreamFrame(expectedFrameHeaderLen(streamID, 10), protocol.Version1) + require.Nil(t, frame.Frame) + require.True(t, hasMore) + frame, _, hasMore = str.popStreamFrame(expectedFrameHeaderLen(streamID, 10)+3, protocol.Version1) + require.True(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Offset: 10, Data: []byte("foo"), DataLenPresent: true}, + frame.Frame, + ) + frame, _, hasMore = str.popStreamFrame(expectedFrameHeaderLen(streamID, 13)+3, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Offset: 13, Data: []byte("baz"), DataLenPresent: true}, + frame.Frame, + ) +} + +func TestSendStreamWriteWithLimit(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, 56) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).Times(4) + _, err := str.Write([]byte("header")) + require.NoError(t, err) + + type result struct { + n int + err error + } + results := make(chan result, 1) + data := make([]byte, 51) + calls := 0 + go func() { + n, err := str.WriteWithLimit(data, func(maxBytes int) int { + calls++ + return min(maxBytes, 25) + }) + results <- result{n: n, err: err} + }() + + synctest.Wait() + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.True(t, hasMore) + require.Equal(t, []byte("header"), frame.Frame.Data) + require.Zero(t, calls) + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.Equal(t, data[:25], frame.Frame.Data) + synctest.Wait() + writeResult := <-results + require.Equal(t, 25, writeResult.n) + require.ErrorIs(t, writeResult.err, ErrWriteLimitReached) + require.Equal(t, 1, calls) + + go func() { + n, err := str.WriteWithLimit(data[25:], func(maxBytes int) int { + calls++ + return maxBytes + }) + results <- result{n: n, err: err} + }() + synctest.Wait() + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.True(t, hasMore) + require.Equal(t, data[25:50], frame.Frame.Data) + require.Equal(t, 2, calls) + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.True(t, hasMore) + require.Equal(t, 2, calls) + + str.updateSendWindow(57) + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.Equal(t, data[50:], frame.Frame.Data) + synctest.Wait() + writeResult = <-results + require.Equal(t, 26, writeResult.n) + require.NoError(t, writeResult.err) + require.Equal(t, 3, calls) + }) +} + +func TestSendStreamTryWriteAll(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + require.NoError(t, str.TryWriteAll(nil)) + require.NoError(t, str.TryWriteAll([]byte{})) + + mockSender.EXPECT().onHasStreamData(streamID, str) + data := []byte("foobar") + require.NoError(t, str.TryWriteAll(data)) + data[0] = 'x' // make sure the data was copied + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foobar"), DataLenPresent: true}, + frame.Frame, + ) +} + +func TestSendStreamTryWriteAllFlowControlBlocked(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, 3) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + require.ErrorIs(t, str.TryWriteAll([]byte("foobar")), ErrWouldBlock) + require.Equal(t, protocol.ByteCount(3), mockFC.SendWindowSize()) + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + + mockSender.EXPECT().onHasStreamData(streamID, str) + require.NoError(t, str.TryWriteAll([]byte("foo"))) + frame, blocked, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, + frame.Frame, + ) + require.Equal(t, &wire.StreamDataBlockedFrame{StreamID: streamID, MaximumStreamData: 3}, blocked) +} + +func TestSendStreamTryWriteAllAfterBufferedWrite(t *testing.T) { + const streamID protocol.StreamID = 42 + + t.Run("enough credit", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, 6) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + n, err := str.Write([]byte("foo")) + require.NoError(t, err) + require.Equal(t, 3, n) + require.NoError(t, str.TryWriteAll([]byte("bar"))) + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foobar"), DataLenPresent: true}, + frame.Frame, + ) + require.Zero(t, mockFC.SendWindowSize()) + }) + + t.Run("not enough credit", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, 5) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + n, err := str.Write([]byte("foo")) + require.NoError(t, err) + require.Equal(t, 3, n) + require.ErrorIs(t, str.TryWriteAll([]byte("bar")), ErrWouldBlock) + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, + frame.Frame, + ) + require.Equal(t, protocol.ByteCount(2), mockFC.SendWindowSize()) + }) +} + +func TestSendStreamWriteAfterTryWriteAll(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + require.NoError(t, str.TryWriteAll([]byte("foo"))) + + errChan := make(chan error, 1) + go func() { + n, err := str.Write([]byte("bar")) + if n != 3 { + errChan <- fmt.Errorf("expected to write 3 bytes, wrote %d", n) + return + } + errChan <- err + }() + synctest.Wait() + select { + case err := <-errChan: + t.Fatalf("Write should not have returned yet: %v", err) + default: + } + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.True(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, + frame.Frame, + ) + + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("Write should have returned") + } + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Offset: 3, Data: []byte("bar"), DataLenPresent: true}, + frame.Frame, + ) + }) +} + +func TestSendStreamLargeTryWriteAll(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + data := make([]byte, 10*protocol.MaxPacketBufferSize) + for i := range data { + data[i] = byte(i) + } + + mockSender.EXPECT().onHasStreamData(streamID, str) + require.NoError(t, str.TryWriteAll(data)) + + var offset protocol.ByteCount + for offset < protocol.ByteCount(len(data)) { + frame, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, offset)+40, protocol.Version1) + require.NotNil(t, frame.Frame) + require.Equal(t, offset, frame.Frame.Offset) + require.Equal(t, data[offset:offset+40], frame.Frame.Data) + offset += 40 + require.Equal(t, offset < protocol.ByteCount(len(data)), hasMore) + } +} + +func TestSendStreamResetFinalSizeIncludesReservedData(t *testing.T) { + const streamID protocol.StreamID = 42 + + for _, tc := range []struct { + name string + reset func(*SendStream) + }{ + { + name: "CancelWrite", + reset: func(str *SendStream) { str.CancelWrite(42) }, + }, + { + name: "STOP_SENDING", + reset: func(str *SendStream) { + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: streamID, ErrorCode: 42}) + }, + }, + } { + t.Run(tc.name, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, 100) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + require.NoError(t, str.TryWriteAll(make([]byte, 100))) + + frame, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 0)+40, protocol.Version1) + require.True(t, hasMore) + require.Equal(t, protocol.ByteCount(40), frame.Frame.DataLen()) + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + tc.reset(str) + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.False(t, hasMore) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 100, ErrorCode: 42}, cf.Frame) + }) + } +} + +func TestSendStreamSetReliableBoundaryAfterTryWriteAll(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, 100) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, true) + + mockSender.EXPECT().onHasStreamData(streamID, str) + require.NoError(t, str.TryWriteAll(make([]byte, 100))) + + frame, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 0)+40, protocol.Version1) + require.True(t, hasMore) + require.Equal(t, protocol.ByteCount(40), frame.Frame.DataLen()) + + str.SetReliableBoundary() + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + str.CancelWrite(42) + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.False(t, hasMore) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 100, ErrorCode: 42, ReliableSize: 100}, cf.Frame) + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.Equal(t, protocol.ByteCount(40), frame.Frame.Offset) + require.Equal(t, protocol.ByteCount(60), frame.Frame.DataLen()) +} + +func TestSendStreamLargeWrites(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 1337 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + data := make([]byte, 5000) + rand.Read(data) + errChan := make(chan error, 1) + go func() { + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write(data) + str.Close() + errChan <- err + }() + + synctest.Wait() + + var offset protocol.ByteCount + const size = 40 + for offset+size < protocol.ByteCount(len(data))-protocol.MaxPacketBufferSize { + frame, _, hasMore := str.popStreamFrame(size+expectedFrameHeaderLen(streamID, offset), protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, hasMore) + require.Equal(t, offset, frame.Frame.Offset) + require.Equal(t, data[offset:offset+size], frame.Frame.Data) + offset += size + require.True(t, mockCtrl.Satisfied()) + } + + // Write should still be blocked, since there's more than protocol.MaxPacketBufferSize left to send + select { + case err := <-errChan: + require.NoError(t, err) + default: + } + + // empty frames are not sent + frame, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, offset), protocol.Version1) + require.Nil(t, frame.Frame) + require.True(t, hasMore) + + mockSender.EXPECT().onHasStreamData(streamID, str) // from the Close call + frame, _, hasMore = str.popStreamFrame(size+expectedFrameHeaderLen(streamID, offset), protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, hasMore) + require.Equal(t, data[offset:offset+size], frame.Frame.Data) + require.Equal(t, offset, frame.Frame.Offset) + offset += size + + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("write should have returned") + } + + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, frame.Frame) + require.False(t, hasMore) + require.Equal(t, data[offset:], frame.Frame.Data) + require.True(t, frame.Frame.Fin) + }) +} + +func TestSendStreamLargeWriteBlocking(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 1337 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write(make([]byte, protocol.MaxPacketBufferSize)) + errChan <- err + }() + + synctest.Wait() + + frame, _, hasMoreData := str.popStreamFrame(expectedFrameHeaderLen(streamID, 0)+3, protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, hasMoreData) + require.Equal(t, []byte("foo"), frame.Frame.Data) + + synctest.Wait() + + select { + case err := <-errChan: + t.Fatalf("write should not have returned yet: %v", err) + default: + } + + frame, _, hasMoreData = str.popStreamFrame(expectedFrameHeaderLen(streamID, 3)+3, protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, hasMoreData) + require.Equal(t, []byte("bar"), frame.Frame.Data) + + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + }) +} + +func TestSendStreamCopyData(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + strWithTimeout := &writerWithTimeout{Writer: str, Timeout: time.Second} + + // for small writes + data := []byte("foobar") + mockSender.EXPECT().onHasStreamData(streamID, str) + _, err := strWithTimeout.Write(data) + require.NoError(t, err) + frame, _, _ := str.popStreamFrame(protocol.MaxPacketBufferSize, protocol.Version1) + data[1] = 'e' // modify the data after it has been written + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foobar"), DataLenPresent: true}, + frame.Frame, + ) +} + +func TestSendStreamDeadlineInThePast(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 42, mockSender, mockFC, false) + + // no data is written when the deadline is in the past + require.NoError(t, str.SetWriteDeadline(time.Now().Add(-time.Second))) + n, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Zero(t, n) + var nerr net.Error + require.ErrorAs(t, err, &nerr) + require.True(t, nerr.Timeout()) + + // data is written when the deadline is in the future + mockSender.EXPECT().onHasStreamData(gomock.Any(), str) + require.NoError(t, str.SetWriteDeadline(time.Now().Add(time.Second))) + n, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + require.Equal(t, 6, n) +} + +func TestSendStreamDeadlineRemoval(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 42, mockSender, mockFC, false) + + deadline := time.Second + require.NoError(t, str.SetWriteDeadline(time.Now().Add(deadline))) + mockSender.EXPECT().onHasStreamData(gomock.Any(), str).Times(2) + + // small writes are written immediately + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + + // large writes might block, and therefore subject to the deadline + errChan := make(chan error, 1) + go func() { + _, err := (&writerWithTimeout{Writer: str, Timeout: 5 * time.Second}).Write(make([]byte, 2000)) + errChan <- err + }() + + synctest.Wait() + + select { + case err := <-errChan: + t.Fatalf("write should not have returned yet: %v", err) + case <-time.After(deadline / 2): + } + + // remove the deadline after a while (but before it expires) + require.NoError(t, str.SetWriteDeadline(time.Time{})) + + select { + case err := <-errChan: + t.Fatalf("write should not have returned yet: %v", err) + case <-time.After(deadline): + } + + // now set the deadline to the past to make Write return immediately + require.NoError(t, str.SetWriteDeadline(time.Now().Add(-time.Second))) + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + default: + } + + frame, _, hasMoreData := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, frame.Frame) + require.False(t, hasMoreData) + require.Equal(t, []byte("foobar"), frame.Frame.Data) + }) +} + +func TestSendStreamDeadlineExtension(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 42, mockSender, mockFC, false) + + deadline := time.Minute + require.NoError(t, str.SetWriteDeadline(time.Now().Add(deadline))) + + mockSender.EXPECT().onHasStreamData(gomock.Any(), str) + errChan := make(chan error, 1) + go func() { + _, err := str.Write(make([]byte, 2000)) + errChan <- err + }() + + synctest.Wait() + + select { + case err := <-errChan: + t.Fatalf("write should not have returned yet: %v", err) + case <-time.After(deadline / 2): + } + + // extend the deadline + start := time.Now() + require.NoError(t, str.SetWriteDeadline(start.Add(deadline))) + + synctest.Wait() + select { + case err := <-errChan: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Equal(t, deadline, time.Since(start)) + case <-time.After(deadline + time.Nanosecond): + t.Fatal("timeout") + } + + frame, _, hasMoreData := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMoreData) + }) +} + +func TestSendStreamClose(t *testing.T) { + const streamID protocol.StreamID = 1234 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + strWithTimeout := &writerWithTimeout{Writer: str, Timeout: time.Second} + + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + _, err := strWithTimeout.Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, str.Close()) + + select { + case <-str.Context().Done(): + default: + t.Fatal("stream context should have been canceled") + } + + frame, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 0)+3, protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Offset: 0, Data: []byte("foo"), DataLenPresent: true}, // no FIN yet + frame.Frame, + ) + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Offset: 3, Fin: true, Data: []byte("bar"), DataLenPresent: true}, + frame.Frame, + ) + require.True(t, mockCtrl.Satisfied()) + + // further calls to Write return an error + _, err = strWithTimeout.Write([]byte("foobar")) + require.ErrorContains(t, err, "write on closed stream 1234") + require.ErrorContains(t, str.TryWriteAll([]byte("foobar")), "write on closed stream 1234") + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + + // further calls to Close don't do anything + require.NoError(t, str.Close()) + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // shutting down has no effect + str.closeForShutdown(errors.New("goodbye")) + _, err = strWithTimeout.Write([]byte("foobar")) + require.ErrorContains(t, err, "write on closed stream 1234") +} + +func TestSendStreamImmediateClose(t *testing.T) { + const streamID protocol.StreamID = 1337 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + mockSender.EXPECT().onHasStreamData(streamID, str) + require.NoError(t, str.Close()) + frame, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 13)+3, protocol.Version1) + require.False(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Fin: true, DataLenPresent: true}, + frame.Frame, + ) +} + +func TestSendStreamFlowControlBlocked(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, 3) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + + frame, blocked, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.True(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, + frame.Frame, + ) + require.Equal(t, &wire.StreamDataBlockedFrame{StreamID: streamID, MaximumStreamData: 3}, blocked) + + frame, blocked, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.Nil(t, blocked) + require.True(t, hasMore) + + _, ok, hasMore := str.getControlFrame(monotime.Now()) + require.False(t, ok) + require.False(t, hasMore) +} + +func TestSendStreamCloseForShutdown(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 1337 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + strWithTimeout := &writerWithTimeout{Writer: str, Timeout: time.Second} + + mockSender.EXPECT().onHasStreamData(streamID, str) + errChan := make(chan error, 1) + go func() { + _, err := strWithTimeout.Write(bytes.Repeat([]byte("foobar"), 1000)) + errChan <- err + }() + + synctest.Wait() + str.closeForShutdown(assert.AnError) + + synctest.Wait() + require.True(t, mockCtrl.Satisfied()) + + select { + case err := <-errChan: + require.ErrorIs(t, err, assert.AnError) + default: + } + + select { + case <-str.Context().Done(): + require.ErrorIs(t, context.Cause(str.Context()), assert.AnError) + default: + t.Fatal("context should be cancelled after closeForShutdown") + } + + // STOP_SENDING frames are ignored + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: streamID, ErrorCode: 1337}) + _, ok, hasMore := str.getControlFrame(monotime.Now()) + require.False(t, ok) + require.False(t, hasMore) + + // future calls to Write should return the error + _, err := strWithTimeout.Write([]byte("foobar")) + require.ErrorIs(t, err, assert.AnError) + + // closing the stream doesn't do anything + require.NoError(t, str.Close()) + + // no STREAM frames popped + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + + // canceling the stream doesn't do anything + str.CancelWrite(1234) + _, err = strWithTimeout.Write([]byte("foobar")) + require.ErrorIs(t, err, assert.AnError) // error unchanged + }) +} + +func TestSendStreamUpdateSendWindow(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, 41) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 42, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(gomock.Any(), str) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + require.True(t, mockCtrl.Satisfied()) + + // no calls to onHasStreamData if the window size wasn't increased + str.updateSendWindow(41) + + mockSender.EXPECT().onHasStreamData(protocol.StreamID(42), str) + str.updateSendWindow(123) +} + +func TestSendStreamCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 42 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + strWithTimeout := &writerWithTimeout{Writer: str, Timeout: time.Second} + + mockSender.EXPECT().onHasStreamData(streamID, str) + _, err := strWithTimeout.Write([]byte("foobar")) + require.NoError(t, err) + frame, _, hasMore := str.popStreamFrame(3+expectedFrameHeaderLen(streamID, 0), protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, hasMore) + require.Equal(t, []byte("foo"), frame.Frame.Data) + require.True(t, mockCtrl.Satisfied()) + + // The stream doesn't support RESET_STREAM_AT. + // Setting the reliable boundary has no effect. + str.SetReliableBoundary() + + wrote := make(chan struct{}) + mockSender.EXPECT().onHasStreamData(streamID, str).Do(func(protocol.StreamID, *SendStream) { close(wrote) }) + errChan := make(chan error, 1) + go func() { + _, err := strWithTimeout.Write(make([]byte, 2000)) + errChan <- err + }() + + synctest.Wait() + + // cancel the stream + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + str.CancelWrite(1234) + require.True(t, mockCtrl.Satisfied()) + + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + // only the "foo" was sent out, so the final size is 3 + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 3, ErrorCode: 1234}, cf.Frame) + require.False(t, hasMore) + + // the context was canceled + select { + case <-str.Context().Done(): + default: + t.Fatal("stream context should have been canceled") + } + require.ErrorIs(t, context.Cause(str.Context()), &StreamError{StreamID: streamID, ErrorCode: 1234, Remote: false}) + + // duplicate calls to CancelWrite don't do anything + str.CancelWrite(1234) + _, ok, _ = str.getControlFrame(monotime.Now()) + require.False(t, ok) + + synctest.Wait() + + // the Write call should return an error + select { + case err := <-errChan: + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1234, Remote: false}) + default: + t.Fatal("write should have returned") + } + + // no data to send + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + + // future calls to Write should return an error + _, err = strWithTimeout.Write([]byte("foo")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1234, Remote: false}) + err = str.TryWriteAll([]byte("foo")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1234, Remote: false}) + frame, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + + // Close has no effect + require.ErrorContains(t, str.Close(), "close called for canceled stream") + frame, _, _ = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + _, err = strWithTimeout.Write([]byte("foobar")) + require.Error(t, err) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1234, Remote: false}) + + // shutting down has no effect + str.closeForShutdown(errors.New("goodbyte")) + _, err = strWithTimeout.Write([]byte("foobar")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1234, Remote: false}) + }) +} + +// It is possible to cancel a stream after it has been closed. +// This is useful if the applications wants to prevent the retransmission of outstanding stream data. +func TestSendStreamCancellationAfterClose(t *testing.T) { + const streamID protocol.StreamID = 1234 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + strWithTimeout := &writerWithTimeout{Writer: str, Timeout: time.Second} + + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + _, err := strWithTimeout.Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, str.Close()) + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + str.CancelWrite(1337) + + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 0, ErrorCode: 1337}, cf.Frame) + require.False(t, hasMore) + + _, err = strWithTimeout.Write([]byte("foobar")) + require.Error(t, err) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: false}) +} + +func TestSendStreamCancellationStreamRetransmission(t *testing.T) { + t.Run("local", func(t *testing.T) { + testSendStreamCancellationStreamRetransmission(t, false) + }) + t.Run("remote", func(t *testing.T) { + testSendStreamCancellationStreamRetransmission(t, true) + }) +} + +func testSendStreamCancellationStreamRetransmission(t *testing.T, remote bool) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + + f1, _, hasMore := str.popStreamFrame(3+expectedFrameHeaderLen(streamID, 0), protocol.Version1) + require.NotNil(t, f1.Frame) + require.True(t, hasMore) + f2, _, hasMore := str.popStreamFrame(3+expectedFrameHeaderLen(streamID, 3), protocol.Version1) + require.NotNil(t, f2.Frame) + require.False(t, hasMore) + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + if remote { + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: streamID, ErrorCode: 1337}) + } else { + str.CancelWrite(1337) + } + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.IsType(t, &wire.ResetStreamFrame{}, cf.Frame) + require.False(t, hasMore) + + // it doesn't matter if the STREAM frames are acked or lost + f1.Handler.OnAcked(f1.Frame) + f2.Handler.OnLost(f2.Frame) + frame, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + require.False(t, hasMore) + // if CancelWrite was called, the stream is completed as soon as the RESET_STREAM frame is acked + if !remote { + mockSender.EXPECT().onStreamCompleted(streamID) + } + cf.Handler.OnAcked(cf.Frame) + + // but if it's a remote cancellation, the application has to consume the error first + if remote { + mockSender.EXPECT().onStreamCompleted(streamID) + _, err := str.Write([]byte("foobar")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + } +} + +func TestSendStreamCancellationResetStreamRetransmission(t *testing.T) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + str.CancelWrite(1337) + + f1, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 0, ErrorCode: 1337}, f1.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // lose the RESET_STREAM frame + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + f1.Handler.OnLost(f1.Frame) + // get the retransmission + f2, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 0, ErrorCode: 1337}, f2.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // acknowledging the RESET_STREAM frame completes the stream + mockSender.EXPECT().onStreamCompleted(streamID) + f2.Handler.OnAcked(f2.Frame) +} + +func TestSendStreamStopSendingAfterWrite(t *testing.T) { + t.Run("complete by Write", func(t *testing.T) { + testSendStreamStopSendingAfterWrite(t, "write") + }) + t.Run("complete by Close", func(t *testing.T) { + testSendStreamStopSendingAfterWrite(t, "close") + }) + t.Run("complete by CancelWrite", func(t *testing.T) { + testSendStreamStopSendingAfterWrite(t, "cancelwrite") + }) +} + +func testSendStreamStopSendingAfterWrite(t *testing.T, completeBy string) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).MaxTimes(2) + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + frame, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, mockCtrl.Satisfied()) + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: streamID, ErrorCode: 1337}) + + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 6, ErrorCode: 1337}, cf.Frame) + require.False(t, hasMore) + + // acknowledging the RESET_STREAM frame doesn't complete the stream, + // since it was neither cancelled nor closed + cf.Handler.OnAcked(cf.Frame) + require.True(t, mockCtrl.Satisfied()) + + mockSender.EXPECT().onStreamCompleted(streamID) + switch completeBy { + case "write": + // calls to Write should return an error + _, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + case "close": + require.ErrorContains(t, str.Close(), "close called for canceled stream") + case "cancelwrite": + str.CancelWrite(1234) + } + // error code and remote flag are unchanged + _, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + frame, _, _ = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + _, ok, _ = str.getControlFrame(monotime.Now()) + require.False(t, ok) +} + +func TestSendStreamStopSendingDuringWrite(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).MaxTimes(2) + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + frame, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, frame.Frame) + require.True(t, mockCtrl.Satisfied()) + + errChan := make(chan error, 1) + go func() { + _, err := str.Write(make([]byte, 2000)) + errChan <- err + }() + + mockSender.EXPECT().onHasStreamControlFrame(streamID, str) + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: streamID, ErrorCode: 1337}) + + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + default: + t.Fatal("write should have returned") + } + + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: streamID, FinalSize: 6, ErrorCode: 1337}, cf.Frame) + require.False(t, hasMore) + + // receiving another STOP_SENDING frame has no effect + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: streamID, ErrorCode: 1234}) + _, ok, hasMore = str.getControlFrame(monotime.Now()) + require.False(t, ok) + require.False(t, hasMore) + + // acknowledging the RESET_STREAM frame completes the stream + mockSender.EXPECT().onStreamCompleted(streamID) + cf.Handler.OnAcked(cf.Frame) + require.True(t, mockCtrl.Satisfied()) + + // calls to Write should return an error + _, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + frame, _, _ = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + + // calls to CancelWrite have no effect + str.CancelWrite(1234) + _, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + // error code and remote flag are unchanged + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + _, ok, _ = str.getControlFrame(monotime.Now()) + require.False(t, ok) + + // Close has no effect + require.ErrorContains(t, str.Close(), "close called for canceled stream") + frame, _, _ = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, frame.Frame) + _, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.Error(t, err) + require.ErrorIs(t, err, &StreamError{StreamID: streamID, ErrorCode: 1337, Remote: true}) + }) +} + +// This test is inherently racy, as it tests a concurrent call to Write() and CancelRead(). +// A single successful run of this test therefore doesn't mean a lot, +// for reliable results it has to be run many times. +func TestSendStreamConcurrentWriteAndCancel(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamControlFrame(gomock.Any(), gomock.Any()).MaxTimes(1) + mockSender.EXPECT().onHasStreamData(streamID, str).MaxTimes(1) + mockSender.EXPECT().onStreamCompleted(streamID).MaxTimes(1) + + errChan := make(chan error, 1) + go func() { + n, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write(make([]byte, 100)) + if n == 0 { + errChan <- nil + return + } + errChan <- err + }() + + done := make(chan struct{}, 2) + go func() { + str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + done <- struct{}{} + }() + go func() { + str.CancelWrite(1234) + done <- struct{}{} + }() + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("write should have returned") + } + + for range 2 { + select { + case <-done: + default: + t.Fatal("timeout waiting for cancel to complete") + } + } + }) +} + +func TestSendStreamRetransmissions(t *testing.T) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + _, err := str.Write([]byte("foo")) + require.NoError(t, err) + + f1, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, + f1.Frame, + ) + require.True(t, mockCtrl.Satisfied()) + + // write some more data + mockSender.EXPECT().onHasStreamData(streamID, str).Times(2) + _, err = (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("bar")) + require.NoError(t, err) + require.NoError(t, str.Close()) + require.True(t, mockCtrl.Satisfied()) + + // lose the frame + mockSender.EXPECT().onHasStreamData(streamID, str) + f1.Handler.OnLost(f1.Frame) + require.True(t, mockCtrl.Satisfied()) + + // when popping a new frame, we first get the retransmission... + f2, _, hasMoreData := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, f2.Frame) + require.True(t, hasMoreData) + require.True(t, mockCtrl.Satisfied()) + + // ... then we get the new data + f3, _, hasMoreData := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, &wire.StreamFrame{StreamID: streamID, Offset: 3, Fin: true, Data: []byte("bar"), DataLenPresent: true}, f3.Frame) + require.False(t, hasMoreData) + require.True(t, mockCtrl.Satisfied()) + + // acknowledge the retransmission... + f2.Handler.OnAcked(f2.Frame) + // ... and the last frame, which concludes this stream + mockSender.EXPECT().onStreamCompleted(streamID) + f3.Handler.OnAcked(f3.Frame) +} + +func TestSendStreamRetransmissionFraming(t *testing.T) { + const streamID protocol.StreamID = 1000 + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str) + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + + f, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, f.Frame) + + // lose the frame + mockSender.EXPECT().onHasStreamData(streamID, str) + f.Handler.OnLost(f.Frame) + + // retransmission doesn't fit + f, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 0), protocol.Version1) + require.Nil(t, f.Frame) + require.True(t, hasMore) + + // split the retransmission + r1, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 0)+3, protocol.Version1) + require.True(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Data: []byte("foo"), DataLenPresent: true}, + r1.Frame, + ) + r2, _, hasMore := str.popStreamFrame(expectedFrameHeaderLen(streamID, 3)+3, protocol.Version1) + require.True(t, hasMore) + // When popping a retransmission, we always claim that there's more data to send. + // We accept that this might be incorrect. + require.True(t, hasMore) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: streamID, Offset: 3, Data: []byte("bar"), DataLenPresent: true}, + r2.Frame, + ) + _, _, hasMore = str.popStreamFrame(expectedFrameHeaderLen(streamID, 3)+3, protocol.Version1) + require.False(t, hasMore) +} + +// This test is kind of an integration test. +// It writes 4 MB of data, and pops STREAM frames that sometimes are and sometimes aren't limited by flow control. +// Half of these STREAM frames are then received and their content saved, while the other half is reported lost +// and has to be retransmitted. +func TestSendStreamRetransmitDataUntilAcknowledged(t *testing.T) { + const streamID protocol.StreamID = 123456 + const dataLen = 1 << 22 // 4 MB + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + mockFC := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, false) + + mockSender.EXPECT().onHasStreamData(streamID, str).AnyTimes() + + data := make([]byte, dataLen) + _, err := rand.Read(data) + require.NoError(t, err) + done := make(chan struct{}) + go func() { + defer close(done) + _, err := str.Write(data) + require.NoError(t, err) + str.Close() + }() + + var completed bool + mockSender.EXPECT().onStreamCompleted(streamID).Do(func(protocol.StreamID) { completed = true }) + + received := make([]byte, dataLen) + var counter int + frameQueue := make([]ackhandler.StreamFrame, 0, 32) + for !completed || len(frameQueue) > 0 { + counter++ + if counter > 1e6 { + t.Fatal("stream should have completed") + } + f, _, _ := str.popStreamFrame(protocol.ByteCount(mrand.IntN(300)+100), protocol.Version1) + var dequeuedFrame bool + if f.Frame != nil { + frameQueue = append(frameQueue, f) + dequeuedFrame = true + } + + // Process one of the queued frames at random. + // This simulates potential reordering. + if len(frameQueue) > 0 && (!dequeuedFrame || len(frameQueue) == cap(frameQueue)) { + idx := mrand.IntN(len(frameQueue)) + f := frameQueue[idx] + // 50%: acknowledge the frame and save the data + // 50%: lose the frame + if mrand.Int()%2 == 0 { + copy(received[f.Frame.Offset:f.Frame.Offset+f.Frame.DataLen()], f.Frame.Data) + f.Handler.OnAcked(f.Frame) + } else { + f.Handler.OnLost(f.Frame) + } + frameQueue = slices.Delete(frameQueue, idx, idx+1) + } + runtime.Gosched() + } + require.Equal(t, data, received) +} + +func TestSendStreamResetStreamAtCancelBeforeSend(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 1337, mockSender, mockFC, true) + + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str).Times(2) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + str.SetReliableBoundary() + _, err = str.Write([]byte("baz")) + require.NoError(t, err) + + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + str.CancelWrite(1337) + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: 1337, FinalSize: 6, ErrorCode: 1337, ReliableSize: 6}, cf.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + f, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("foobar"), DataLenPresent: true}, + f.Frame, + ) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // Lose the frame. + // Since it's before the reliable size, we should get a retransmission. + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str) + f.Handler.OnLost(f.Frame) + require.True(t, mockCtrl.Satisfied()) + + retransmission, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("foobar"), DataLenPresent: true}, + retransmission.Frame, + ) + require.True(t, hasMore) // hasMore is always true when dequeuing a retransmission + require.True(t, mockCtrl.Satisfied()) + f, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, f.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // acknowledging the RESET_STREAM_AT and the retransmission completes the stream + cf.Handler.OnAcked(cf.Frame) + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(1337)) + retransmission.Handler.OnAcked(retransmission.Frame) +} + +func TestSendStreamResetStreamAtCancelAfterSend(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 1337, mockSender, mockFC, true) + + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str).Times(2) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + str.SetReliableBoundary() + _, err = str.Write([]byte("baz")) + require.NoError(t, err) + + f, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("foobarbaz"), DataLenPresent: true}, + f.Frame, + ) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + str.CancelWrite(42) + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: 1337, FinalSize: 9, ErrorCode: 42, ReliableSize: 6}, cf.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + cf.Handler.OnAcked(cf.Frame) + // lose the STREAM frame + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str) + f.Handler.OnLost(f.Frame) + // only the first 6 bytes need to be retransmitted + retransmission1, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("foobar"), DataLenPresent: true}, + retransmission1.Frame, + ) + require.True(t, hasMore) // hasMore is always true when dequeuing a retransmission + require.True(t, mockCtrl.Satisfied()) + f, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, f.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // lose the retransmission as well + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str) + retransmission1.Handler.OnLost(retransmission1.Frame) + retransmission2, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("foobar"), DataLenPresent: true}, + retransmission2.Frame, + ) + require.True(t, hasMore) // hasMore is always true when dequeuing a retransmission + require.True(t, mockCtrl.Satisfied()) + f, _, hasMore = str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, f.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // acknowledge the 2nd retransmission + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(1337)) + retransmission2.Handler.OnAcked(retransmission2.Frame) +} + +func TestSendStreamResetStreamAtRetransmissions(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 1337, mockSender, mockFC, true) + + // f1: lorem + // f2: ipsumdolor (reliable offset: right after the "ipsum") + // f3: sit + // f4: amet + // sitting in the write buffer: consectetur (but not popped) + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str).AnyTimes() + _, err := str.Write([]byte("lorem")) + require.NoError(t, err) + f1, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("lorem"), DataLenPresent: true}, + f1.Frame, + ) + _, err = str.Write([]byte("ipsum")) + require.NoError(t, err) + str.SetReliableBoundary() + _, err = str.Write([]byte("dolor")) + require.NoError(t, err) + f2, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Offset: 5, Data: []byte("ipsumdolor"), DataLenPresent: true}, + f2.Frame, + ) + _, err = str.Write([]byte("sit")) + require.NoError(t, err) + f3, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Offset: 15, Data: []byte("sit"), DataLenPresent: true}, + f3.Frame, + ) + _, err = str.Write([]byte("amet")) + require.NoError(t, err) + f4, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Offset: 18, Data: []byte("amet"), DataLenPresent: true}, + f4.Frame, + ) + _, err = str.Write([]byte("consectetur")) + require.NoError(t, err) + + // lose the frames, in no particular order + f2.Handler.OnLost(f2.Frame) + f1.Handler.OnLost(f1.Frame) + f3.Handler.OnLost(f3.Frame) + // f4 is lost at a later point + + // Now cancel the stream. + // We expect f1 and the first half of f2 to be retransmitted, + // but f3 and the data in the buffer should not. + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + str.CancelWrite(42) + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: 1337, FinalSize: 22, ErrorCode: 42, ReliableSize: 10}, cf.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + cf.Handler.OnAcked(cf.Frame) + + // // the retransmission of f1 should be truncated to 6 bytes + r1, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Offset: 5, Data: []byte("ipsum"), DataLenPresent: true}, + r1.Frame, + ) + require.True(t, hasMore) + r2, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.EqualExportedValues(t, + &wire.StreamFrame{StreamID: 1337, Data: []byte("lorem"), DataLenPresent: true}, + r2.Frame, + ) + require.True(t, hasMore) // hasMore is always true when dequeuing a retransmission + require.True(t, mockCtrl.Satisfied()) + r3, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Nil(t, r3.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + r1.Handler.OnAcked(r1.Frame) + r2.Handler.OnAcked(r2.Frame) + require.True(t, mockCtrl.Satisfied()) + + // the stream is only completed once f4 is lost + // it's beyond the reliable size, so it's not retransmitted + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(1337)) + f4.Handler.OnLost(f4.Frame) +} + +func TestSendStreamResetStreamAtStopSendingBeforeCancelation(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 1337, mockSender, mockFC, true) + + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str).Times(2) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + str.SetReliableBoundary() + _, err = str.Write([]byte("baz")) + require.NoError(t, err) + + // send out a STREAM frame with all the data written so far + f, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Equal(t, protocol.ByteCount(9), f.Frame.DataLen()) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: 1337, ErrorCode: 42}) + cf, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + // Since the peer reset the stream, the resulting RESET_STREAM frame has a reliable size of 0 + require.Equal(t, &wire.ResetStreamFrame{StreamID: 1337, FinalSize: 9, ErrorCode: 42, ReliableSize: 0}, cf.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // calling CancelWrite doesn't cause any more frames to be enqueued + str.CancelWrite(1234) + + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(1337)) + cf.Handler.OnAcked(cf.Frame) +} + +func TestSendStreamResetStreamAtStopSendingAfterCancelation(t *testing.T) { + t.Run("RESET_STREAM_AT lost", func(t *testing.T) { + testSendStreamResetStreamAtStopSendingAfterCancelation(t, true) + }) + t.Run("RESET_STREAM_AT acknowledged", func(t *testing.T) { + testSendStreamResetStreamAtStopSendingAfterCancelation(t, false) + }) +} + +func testSendStreamResetStreamAtStopSendingAfterCancelation(t *testing.T, loseResetStreamAt bool) { + mockCtrl := gomock.NewController(t) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + mockSender := NewMockStreamSender(mockCtrl) + str := newSendStream(context.Background(), 1337, mockSender, mockFC, true) + + mockSender.EXPECT().onHasStreamData(protocol.StreamID(1337), str).Times(2) + _, err := str.Write([]byte("foobar")) + require.NoError(t, err) + str.SetReliableBoundary() + _, err = str.Write([]byte("baz")) + require.NoError(t, err) + + // send out a STREAM frame with all the data written so far + f, _, hasMore := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.Equal(t, protocol.ByteCount(9), f.Frame.DataLen()) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + // Canceling the stream results in a RESET_STREAM_AT frame. + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + str.CancelWrite(42) + cf1, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, &wire.ResetStreamFrame{StreamID: 1337, FinalSize: 9, ErrorCode: 42, ReliableSize: 6}, cf1.Frame) + require.False(t, hasMore) + + // Receiving a STOP_SENDING frame results in a RESET_STREAM frame, + // effectively reducing the reliable size to 0. + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + str.handleStopSendingFrame(&wire.StopSendingFrame{StreamID: 1337, ErrorCode: 1234}) + cf2, ok, hasMore := str.getControlFrame(monotime.Now()) + require.True(t, ok) + // Since the peer reset the stream, the resulting RESET_STREAM frame has a reliable size of 0. + // The error code is still the one used for the CancelWrite call. + require.Equal(t, &wire.ResetStreamFrame{StreamID: 1337, FinalSize: 9, ErrorCode: 42, ReliableSize: 0}, cf2.Frame) + require.False(t, hasMore) + require.True(t, mockCtrl.Satisfied()) + + if loseResetStreamAt { + // losing the RESET_STREAM_AT frame does nothing + cf1.Handler.OnLost(cf1.Frame) + } else { + // receiving an acknowledgment for the RESET_STREAM_AT frame does nothing either: + // the RESET_STREAM frame still needs to be transmitted reliably + cf1.Handler.OnAcked(cf1.Frame) + } + _, ok, _ = str.getControlFrame(monotime.Now()) + require.False(t, ok) + + // but when the RESET_STREAM frame is lost, it needs to be retransmitted + mockSender.EXPECT().onHasStreamControlFrame(protocol.StreamID(1337), str) + cf2.Handler.OnLost(cf2.Frame) + cf3, ok, _ := str.getControlFrame(monotime.Now()) + require.True(t, ok) + require.Equal(t, cf2, cf3) + + mockSender.EXPECT().onStreamCompleted(protocol.StreamID(1337)) + cf3.Handler.OnAcked(cf3.Frame) +} + +func TestSendStreamResetStreamAtRandomized(t *testing.T) { + const streamID protocol.StreamID = 123456 + const dataLen = 8 << 10 + reliableOffset := 1 + mrand.IntN(dataLen*3/4) + t.Logf("reliable offset: %d", reliableOffset) + + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + mockFC := newTestStreamFlowControllerWithSendWindow(42, protocol.MaxByteCount) + str := newSendStream(context.Background(), streamID, mockSender, mockFC, true) + + mockSender.EXPECT().onHasStreamData(streamID, str).AnyTimes() + mockSender.EXPECT().onHasStreamControlFrame(streamID, str).AnyTimes() + + data := make([]byte, dataLen) + _, err := rand.Read(data) + require.NoError(t, err) + errChan := make(chan error, 1) + go func() { + b := data + var offset int + for len(b) > 0 { + m := mrand.IntN(1024) + if offset < reliableOffset { + m = min(m, reliableOffset-offset) + } + n, err := str.Write(b[:min(m, len(b))]) + if err != nil { + errChan <- err + return + } + offset += n + if offset <= reliableOffset { + str.SetReliableBoundary() + } + b = b[n:] + } + str.CancelWrite(1234) + errChan <- nil + }() + + var completed bool + mockSender.EXPECT().onStreamCompleted(streamID).Do(func(protocol.StreamID) { completed = true }) + + received := make([]byte, dataLen) + var highestOffset int + var receivedResetStreamAt bool + var counter int + frameQueue := make([]any, 0, 10) + for !completed || len(frameQueue) > 0 { + counter++ + if counter > 1e6 { + t.Fatal("stream should have completed") + } + var dequeuedFrame bool + cf, ok, _ := str.getControlFrame(monotime.Now()) + if ok { + dequeuedFrame = true + frameQueue = append(frameQueue, cf) + receivedResetStreamAt = true + require.Equal(t, protocol.ByteCount(reliableOffset), cf.Frame.(*wire.ResetStreamFrame).ReliableSize) + } else { + f, _, _ := str.popStreamFrame(protocol.ByteCount(mrand.IntN(300)+100), protocol.Version1) + if f.Frame != nil { + // make sure that only retransmissions are sent once the RESET_STREAM_AT frame is sent + if receivedResetStreamAt { + require.LessOrEqualf(t, + f.Frame.Offset+f.Frame.DataLen(), + protocol.ByteCount(reliableOffset), + "STREAM frame past reliable offset after RESET_STREAM_AT (offset: %d, data length: %d)", + f.Frame.Offset, f.Frame.DataLen(), + ) + } + dequeuedFrame = true + frameQueue = append(frameQueue, f) + } + } + + if len(frameQueue) > 0 && (!dequeuedFrame || len(frameQueue) == cap(frameQueue)) { + idx := mrand.IntN(len(frameQueue)) + switch f := frameQueue[idx].(type) { + case ackhandler.Frame: + // 50%: acknowledge the frame + // 50%: lose the frame + if mrand.Int()%2 == 0 { + f.Handler.OnLost(f.Frame) + } else { + f.Handler.OnAcked(f.Frame) + } + case ackhandler.StreamFrame: + sf := f.Frame + // 50%: acknowledge the frame and save the data + // 50%: lose the frame + if mrand.Int()%2 == 0 { + f.Handler.OnLost(f.Frame) + } else { + highestOffset = max(highestOffset, int(sf.Offset+sf.DataLen())) + copy(received[sf.Offset:sf.Offset+sf.DataLen()], sf.Data) + f.Handler.OnAcked(f.Frame) + } + default: + t.Fatalf("unexpected frame type: %T", f) + } + frameQueue = slices.Delete(frameQueue, idx, idx+1) + } + runtime.Gosched() + } + + t.Logf("highest received offset: %d", highestOffset) + require.GreaterOrEqual(t, highestOffset, reliableOffset) + require.Equal(t, data[:reliableOffset], received[:reliableOffset]) +} diff --git a/third_party/quic-go/server.go b/third_party/quic-go/server.go new file mode 100644 index 0000000..442c34b --- /dev/null +++ b/third_party/quic-go/server.go @@ -0,0 +1,1128 @@ +package quic + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "sync" + "time" + + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +// ErrServerClosed is returned by the [Listener] or [EarlyListener]'s Accept method after a call to Close. +var ErrServerClosed = errServerClosed{} + +type errServerClosed struct{} + +func (errServerClosed) Error() string { return "quic: server closed" } +func (errServerClosed) Unwrap() error { return net.ErrClosed } + +// packetHandler handles packets +type packetHandler interface { + handlePacket(receivedPacket) + destroy(error) + closeWithTransportError(qerr.TransportErrorCode) +} + +type zeroRTTQueue struct { + packets []receivedPacket + expiration monotime.Time +} + +type rejectedPacket struct { + receivedPacket + hdr *wire.Header +} + +// A Listener of QUIC +type baseServer struct { + tr *packetHandlerMap + disableVersionNegotiation bool + acceptEarlyConns bool + + tlsConf *tls.Config + config *Config + + conn rawConn + + tokenGenerator *handshake.TokenGenerator + maxTokenAge time.Duration + + connIDGenerator ConnectionIDGenerator + statelessResetter *statelessResetter + onClose func() + + receivedPackets chan receivedPacket + + nextZeroRTTCleanup monotime.Time + zeroRTTQueues map[protocol.ConnectionID]*zeroRTTQueue // only initialized if acceptEarlyConns == true + + connContext func(context.Context, *ClientInfo) (context.Context, error) + + // set as a member, so they can be set in the tests + newConn func( + context.Context, + context.CancelCauseFunc, + sendConn, + connRunner, + protocol.ConnectionID, /* original dest connection ID */ + *protocol.ConnectionID, /* retry src connection ID */ + protocol.ConnectionID, /* client dest connection ID */ + protocol.ConnectionID, /* destination connection ID */ + protocol.ConnectionID, /* source connection ID */ + ConnectionIDGenerator, + *statelessResetter, + *Config, + *tls.Config, + *handshake.TokenGenerator, + bool, /* client address validated by an address validation token */ + time.Duration, + qlogwriter.Trace, + utils.Logger, + protocol.Version, + ) *wrappedConn + + closeMx sync.Mutex + // errorChan is closed when Close is called. This has two effects: + // 1. it cancels handshakes that are still in flight (using CONNECTION_REFUSED) errors + // 2. it stops handling of packets passed to this server + errorChan chan struct{} + // acceptChan is closed when Close returns. + // This only happens once all handshake in flight have either completed and canceled. + // Calls to Accept will first drain the queue of connections that have completed the handshake, + // and then return ErrServerClosed. + stopAccepting chan struct{} + closeErr error + running chan struct{} // closed as soon as run() returns + + versionNegotiationQueue chan receivedPacket + invalidTokenQueue chan rejectedPacket + connectionRefusedQueue chan rejectedPacket + retryQueue chan rejectedPacket + handshakingCount sync.WaitGroup + + verifySourceAddress func(net.Addr) bool + + connQueue chan *Conn + + qlogger qlogwriter.Recorder + + logger utils.Logger +} + +// A Listener listens for incoming QUIC connections. +// It returns connections once the handshake has completed. +type Listener struct { + baseServer *baseServer +} + +// Accept returns new connections. It should be called in a loop. +func (l *Listener) Accept(ctx context.Context) (*Conn, error) { + return l.baseServer.Accept(ctx) +} + +// Close closes the listener. +// Accept will return [ErrServerClosed] as soon as all connections in the accept queue have been accepted. +// QUIC handshakes that are still in flight will be rejected with a CONNECTION_REFUSED error. +// Already established (accepted) connections will be unaffected. +func (l *Listener) Close() error { + return l.baseServer.Close() +} + +// Addr returns the local network address that the server is listening on. +func (l *Listener) Addr() net.Addr { + return l.baseServer.Addr() +} + +// An EarlyListener listens for incoming QUIC connections, and returns them before the handshake completes. +// For connections that don't use 0-RTT, this allows the server to send 0.5-RTT data. +// This data is encrypted with forward-secure keys, however, the client's identity has not yet been verified. +// For connection using 0-RTT, this allows the server to accept and respond to streams that the client opened in the +// 0-RTT data it sent. Note that at this point during the handshake, the live-ness of the +// client has not yet been confirmed, and the 0-RTT data could have been replayed by an attacker. +type EarlyListener struct { + baseServer *baseServer +} + +// Accept returns a new connections. It should be called in a loop. +func (l *EarlyListener) Accept(ctx context.Context) (*Conn, error) { + conn, err := l.baseServer.accept(ctx) + if err != nil { + return nil, err + } + return conn, nil +} + +// Close closes the listener. +// Accept will return [ErrServerClosed] as soon as all connections in the accept queue have been accepted. +// Early connections that are still in flight will be rejected with a CONNECTION_REFUSED error. +// Already established (accepted) connections will be unaffected. +func (l *EarlyListener) Close() error { + return l.baseServer.Close() +} + +// Addr returns the local network addr that the server is listening on. +func (l *EarlyListener) Addr() net.Addr { + return l.baseServer.Addr() +} + +// ListenAddr creates a QUIC server listening on a given address. +// See [Listen] for more details. +func ListenAddr(addr string, tlsConf *tls.Config, config *Config) (*Listener, error) { + conn, err := listenUDP(addr) + if err != nil { + return nil, err + } + return (&Transport{ + Conn: conn, + createdConn: true, + isSingleUse: true, + }).Listen(tlsConf, config) +} + +// ListenAddrEarly works like [ListenAddr], but it returns connections before the handshake completes. +func ListenAddrEarly(addr string, tlsConf *tls.Config, config *Config) (*EarlyListener, error) { + conn, err := listenUDP(addr) + if err != nil { + return nil, err + } + return (&Transport{ + Conn: conn, + createdConn: true, + isSingleUse: true, + }).ListenEarly(tlsConf, config) +} + +func listenUDP(addr string) (*net.UDPConn, error) { + udpAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return nil, err + } + return net.ListenUDP("udp", udpAddr) +} + +// Listen listens for QUIC connections on a given net.PacketConn. +// If the PacketConn satisfies the [OOBCapablePacketConn] interface (as a [net.UDPConn] does), +// ECN and packet info support will be enabled. In this case, ReadMsgUDP and WriteMsgUDP +// will be used instead of ReadFrom and WriteTo to read/write packets. +// A single net.PacketConn can only be used for a single call to Listen. +// +// The tls.Config must not be nil and must contain a certificate configuration. +// Furthermore, it must define an application control (using [NextProtos]). +// The quic.Config may be nil, in that case the default values will be used. +// +// This is a convenience function. More advanced use cases should instantiate a [Transport], +// which offers configuration options for a more fine-grained control of the connection establishment, +// including reusing the underlying UDP socket for outgoing QUIC connections. +// When closing a listener created with Listen, all established QUIC connections will be closed immediately. +func Listen(conn net.PacketConn, tlsConf *tls.Config, config *Config) (*Listener, error) { + tr := &Transport{Conn: conn, isSingleUse: true} + return tr.Listen(tlsConf, config) +} + +// ListenEarly works like [Listen], but it returns connections before the handshake completes. +func ListenEarly(conn net.PacketConn, tlsConf *tls.Config, config *Config) (*EarlyListener, error) { + tr := &Transport{Conn: conn, isSingleUse: true} + return tr.ListenEarly(tlsConf, config) +} + +func newServer( + conn rawConn, + tr *packetHandlerMap, + connIDGenerator ConnectionIDGenerator, + statelessResetter *statelessResetter, + connContext func(context.Context, *ClientInfo) (context.Context, error), + tlsConf *tls.Config, + config *Config, + qlogger qlogwriter.Recorder, + onClose func(), + tokenGeneratorKey TokenGeneratorKey, + maxTokenAge time.Duration, + verifySourceAddress func(net.Addr) bool, + disableVersionNegotiation bool, + acceptEarly bool, +) *baseServer { + s := &baseServer{ + conn: conn, + connContext: connContext, + tr: tr, + tlsConf: tlsConf, + config: config, + tokenGenerator: handshake.NewTokenGenerator(tokenGeneratorKey), + maxTokenAge: maxTokenAge, + verifySourceAddress: verifySourceAddress, + connIDGenerator: connIDGenerator, + statelessResetter: statelessResetter, + connQueue: make(chan *Conn, protocol.MaxAcceptQueueSize), + errorChan: make(chan struct{}), + stopAccepting: make(chan struct{}), + running: make(chan struct{}), + receivedPackets: make(chan receivedPacket, protocol.MaxServerUnprocessedPackets), + versionNegotiationQueue: make(chan receivedPacket, 4), + invalidTokenQueue: make(chan rejectedPacket, 4), + connectionRefusedQueue: make(chan rejectedPacket, 4), + retryQueue: make(chan rejectedPacket, 8), + newConn: newConnection, + qlogger: qlogger, + logger: utils.DefaultLogger.WithPrefix("server"), + acceptEarlyConns: acceptEarly, + disableVersionNegotiation: disableVersionNegotiation, + onClose: onClose, + } + if acceptEarly { + s.zeroRTTQueues = map[protocol.ConnectionID]*zeroRTTQueue{} + } + go s.run() + go s.runSendQueue() + s.logger.Debugf("Listening for %s connections on %s", conn.LocalAddr().Network(), conn.LocalAddr().String()) + return s +} + +func (s *baseServer) run() { + defer close(s.running) + for { + select { + case <-s.errorChan: + return + default: + } + select { + case <-s.errorChan: + return + case p := <-s.receivedPackets: + if bufferStillInUse := s.handlePacketImpl(p); !bufferStillInUse { + p.buffer.Release() + } + } + } +} + +func (s *baseServer) runSendQueue() { + for { + select { + case <-s.running: + return + case p := <-s.versionNegotiationQueue: + s.maybeSendVersionNegotiationPacket(p) + case p := <-s.invalidTokenQueue: + s.maybeSendInvalidToken(p) + case p := <-s.connectionRefusedQueue: + s.sendConnectionRefused(p) + case p := <-s.retryQueue: + s.sendRetry(p) + } + } +} + +// Accept returns connections that already completed the handshake. +// It is only valid if acceptEarlyConns is false. +func (s *baseServer) Accept(ctx context.Context) (*Conn, error) { + return s.accept(ctx) +} + +func (s *baseServer) accept(ctx context.Context) (*Conn, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case conn := <-s.connQueue: + return conn, nil + case <-s.stopAccepting: + // first drain the queue + select { + case conn := <-s.connQueue: + return conn, nil + default: + } + return nil, s.closeErr + } +} + +func (s *baseServer) Close() error { + s.close(ErrServerClosed, false) + return nil +} + +// close closes the server. The Transport mutex must not be held while calling this method. +// This method closes any handshaking connections which requires the tranpsort mutex. +func (s *baseServer) close(e error, transportClose bool) { + s.closeMx.Lock() + if s.closeErr != nil { + s.closeMx.Unlock() + return + } + s.closeErr = e + close(s.errorChan) + <-s.running + s.closeMx.Unlock() + + if !transportClose { + s.onClose() + } + + // wait until all handshakes in flight have terminated + s.handshakingCount.Wait() + close(s.stopAccepting) + + if transportClose { + // if the transport is closing, drain the connQueue. All connections in the queue + // will be closed by the transport. + for { + select { + case <-s.connQueue: + default: + return + } + } + } +} + +// Addr returns the server's network address +func (s *baseServer) Addr() net.Addr { + return s.conn.LocalAddr() +} + +func (s *baseServer) handlePacket(p receivedPacket) { + select { + case s.receivedPackets <- p: + case <-s.errorChan: + return + default: + s.logger.Debugf("Dropping packet from %s (%d bytes). Server receive queue full.", p.remoteAddr, p.Size()) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + } +} + +func (s *baseServer) handlePacketImpl(p receivedPacket) bool /* is the buffer still in use? */ { + if !s.nextZeroRTTCleanup.IsZero() && p.rcvTime.After(s.nextZeroRTTCleanup) { + defer s.cleanupZeroRTTQueues(p.rcvTime) + } + + if wire.IsVersionNegotiationPacket(p.data) { + s.logger.Debugf("Dropping Version Negotiation packet.") + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false + } + // Short header packets should never end up here in the first place + if !wire.IsLongHeaderPacket(p.data[0]) { + panic(fmt.Sprintf("misrouted packet: %#v", p.data)) + } + v, err := wire.ParseVersion(p.data) + // drop the packet if we failed to parse the protocol version + if err != nil { + s.logger.Debugf("Dropping a packet with an unknown version") + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false + } + // send a Version Negotiation Packet if the client is speaking a different protocol version + if !protocol.IsSupportedVersion(s.config.Versions, v) { + if s.disableVersionNegotiation { + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{Version: v}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedVersion, + }) + } + return false + } + + if p.Size() < protocol.MinUnknownVersionPacketSize { + s.logger.Debugf("Dropping a packet with an unsupported version number %d that is too small (%d bytes)", v, p.Size()) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{Version: v}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false + } + return s.enqueueVersionNegotiationPacket(p) + } + + if wire.Is0RTTPacket(p.data) { + if !s.acceptEarlyConns { + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false + } + return s.handle0RTTPacket(p) + } + + // If we're creating a new connection, the packet will be passed to the connection. + // The header will then be parsed again. + hdr, _, _, err := wire.ParsePacket(p.data) + if err != nil { + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + s.logger.Debugf("Error parsing packet: %s", err) + return false + } + if hdr.Type == protocol.PacketTypeInitial && p.Size() < protocol.MinInitialPacketSize { + s.logger.Debugf("Dropping a packet that is too small to be a valid Initial (%d bytes)", p.Size()) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + Version: v, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false + } + + if hdr.Type != protocol.PacketTypeInitial { + // Drop long header packets. + // There's little point in sending a Stateless Reset, since the client + // might not have received the token yet. + s.logger.Debugf("Dropping long header packet of type %s (%d bytes)", hdr.Type, len(p.data)) + if s.qlogger != nil { + var pt qlog.PacketType + switch hdr.Type { + case protocol.PacketTypeInitial: + pt = qlog.PacketTypeInitial + case protocol.PacketTypeHandshake: + pt = qlog.PacketTypeHandshake + case protocol.PacketType0RTT: + pt = qlog.PacketType0RTT + case protocol.PacketTypeRetry: + pt = qlog.PacketTypeRetry + } + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: pt, + PacketNumber: protocol.InvalidPacketNumber, + Version: v, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return false + } + + s.logger.Debugf("<- Received Initial packet.") + + if err := s.handleInitialImpl(p, hdr); err != nil { + s.logger.Errorf("Error occurred handling initial packet: %s", err) + } + // Don't put the packet buffer back. + // handleInitialImpl deals with the buffer. + return true +} + +func (s *baseServer) handle0RTTPacket(p receivedPacket) bool { + connID, err := wire.ParseConnectionID(p.data, 0) + if err != nil { + if s.qlogger != nil { + v, _ := wire.ParseVersion(p.data) + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: v, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + return false + } + + // check again if we might have a connection now + if handler, ok := s.tr.Get(connID); ok { + handler.handlePacket(p) + return true + } + + if q, ok := s.zeroRTTQueues[connID]; ok { + if len(q.packets) >= protocol.Max0RTTQueueLen { + if s.qlogger != nil { + v, _ := wire.ParseVersion(p.data) + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: v, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + return false + } + q.packets = append(q.packets, p) + return true + } + + if len(s.zeroRTTQueues) >= protocol.Max0RTTQueues { + if s.qlogger != nil { + v, _ := wire.ParseVersion(p.data) + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: v, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + return false + } + queue := &zeroRTTQueue{packets: make([]receivedPacket, 1, 8)} + queue.packets[0] = p + expiration := p.rcvTime.Add(protocol.Max0RTTQueueingDuration) + queue.expiration = expiration + if s.nextZeroRTTCleanup.IsZero() || s.nextZeroRTTCleanup.After(expiration) { + s.nextZeroRTTCleanup = expiration + } + s.zeroRTTQueues[connID] = queue + return true +} + +func (s *baseServer) cleanupZeroRTTQueues(now monotime.Time) { + // Iterate over all queues to find those that are expired. + // This is ok since we're placing a pretty low limit on the number of queues. + var nextCleanup monotime.Time + for connID, q := range s.zeroRTTQueues { + if q.expiration.After(now) { + if nextCleanup.IsZero() || nextCleanup.After(q.expiration) { + nextCleanup = q.expiration + } + continue + } + for _, p := range q.packets { + if s.qlogger != nil { + v, _ := wire.ParseVersion(p.data) + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: v, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + p.buffer.Release() + } + delete(s.zeroRTTQueues, connID) + if s.logger.Debug() { + s.logger.Debugf("Removing 0-RTT queue for %s.", connID) + } + } + s.nextZeroRTTCleanup = nextCleanup +} + +// validateToken returns false if: +// - address is invalid +// - token is expired +// - token is null +func (s *baseServer) validateToken(token *handshake.Token, addr net.Addr) bool { + if token == nil { + return false + } + if !token.ValidateRemoteAddr(addr) { + return false + } + if !token.IsRetryToken && time.Since(token.SentTime) > s.maxTokenAge { + return false + } + if token.IsRetryToken && time.Since(token.SentTime) > s.config.maxRetryTokenAge() { + return false + } + return true +} + +func (s *baseServer) handleInitialImpl(p receivedPacket, hdr *wire.Header) error { + if len(hdr.Token) == 0 && hdr.DestConnectionID.Len() < protocol.MinConnectionIDLenInitial { + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + p.buffer.Release() + return errors.New("too short connection ID") + } + + // The server queues packets for a while, and we might already have established a connection by now. + // This results in a second check in the connection map. + // That's ok since it's not the hot path (it's only taken by some Initial and 0-RTT packets). + if handler, ok := s.tr.Get(hdr.DestConnectionID); ok { + handler.handlePacket(p) + return nil + } + + var ( + token *handshake.Token + retrySrcConnID *protocol.ConnectionID + clientAddrVerified bool + ) + origDestConnID := hdr.DestConnectionID + if len(hdr.Token) > 0 { + tok, err := s.tokenGenerator.DecodeToken(hdr.Token) + if err == nil { + if tok.IsRetryToken { + origDestConnID = tok.OriginalDestConnectionID + retrySrcConnID = &tok.RetrySrcConnectionID + } + token = tok + } + } + if token != nil { + clientAddrVerified = s.validateToken(token, p.remoteAddr) + if !clientAddrVerified { + // For invalid and expired non-retry tokens, we don't send an INVALID_TOKEN error. + // We just ignore them, and act as if there was no token on this packet at all. + // This also means we might send a Retry later. + if !token.IsRetryToken { + token = nil + } else { + // For Retry tokens, we send an INVALID_ERROR if + // * the token is too old, or + // * the token is invalid, in case of a retry token. + select { + case s.invalidTokenQueue <- rejectedPacket{receivedPacket: p, hdr: hdr}: + default: + // drop packet if we can't send out the INVALID_TOKEN packets fast enough + p.buffer.Release() + } + return nil + } + } + } + + if token == nil && s.verifySourceAddress != nil && s.verifySourceAddress(p.remoteAddr) { + // Retry invalidates all 0-RTT packets sent. + delete(s.zeroRTTQueues, hdr.DestConnectionID) + select { + case s.retryQueue <- rejectedPacket{receivedPacket: p, hdr: hdr}: + default: + // drop packet if we can't send out Retry packets fast enough + p.buffer.Release() + } + return nil + } + + // restore RTT from token + var rtt time.Duration + if token != nil && !token.IsRetryToken { + rtt = token.RTT + } + + config := s.config + clientInfo := &ClientInfo{ + RemoteAddr: p.remoteAddr, + AddrVerified: clientAddrVerified, + } + if s.config.GetConfigForClient != nil { + conf, err := s.config.GetConfigForClient(clientInfo) + if err != nil { + s.logger.Debugf("Rejecting new connection due to GetConfigForClient callback") + s.refuseNewConn(p, hdr) + return nil + } + config = populateConfig(conf) + } + + var conn *wrappedConn + var cancel context.CancelCauseFunc + ctx, cancel1 := context.WithCancelCause(context.Background()) + if s.connContext != nil { + var err error + ctx, err = s.connContext(ctx, clientInfo) + if err != nil { + cancel1(err) + s.logger.Debugf("Rejecting new connection due to ConnContext callback: %s", err) + s.refuseNewConn(p, hdr) + return nil + } + if ctx == nil { + panic("quic: ConnContext returned nil") + } + // There's no guarantee that the application returns a context + // that's derived from the context we passed into ConnContext. + // We need to make sure that both contexts are cancelled. + var cancel2 context.CancelCauseFunc + ctx, cancel2 = context.WithCancelCause(ctx) + cancel = func(cause error) { + cancel1(cause) + cancel2(cause) + } + } else { + cancel = cancel1 + } + var qlogTrace qlogwriter.Trace + if config.Tracer != nil { + // Use the same connection ID that is passed to the client's GetLogWriter callback. + connID := hdr.DestConnectionID + if origDestConnID.Len() > 0 { + connID = origDestConnID + } + qlogTrace = config.Tracer(ctx, false, connID) + } + connID, err := s.connIDGenerator.GenerateConnectionID() + if err != nil { + // ConnContext may already have reserved application capacity. No + // connection object exists yet to own cancellation or the packet buffer, + // so unwind both explicitly on this rare random-source failure. + cancel(err) + if qlogTrace != nil { + _ = qlogTrace.AddProducer().Close() + } + delete(s.zeroRTTQueues, hdr.DestConnectionID) + p.buffer.Release() + return err + } + s.logger.Debugf("Changing connection ID to %s.", connID) + conn = s.newConn( + ctx, + cancel, + newSendConn(s.conn, p.remoteAddr, p.info, s.logger), + s.tr, + origDestConnID, + retrySrcConnID, + hdr.DestConnectionID, + hdr.SrcConnectionID, + connID, + s.connIDGenerator, + s.statelessResetter, + config, + s.tlsConf, + s.tokenGenerator, + clientAddrVerified, + rtt, + qlogTrace, + s.logger, + hdr.Version, + ) + conn.handlePacket(p) + // Adding the connection will fail if the client's chosen Destination Connection ID is already in use. + // This is very unlikely: Even if an attacker chooses a connection ID that's already in use, + // under normal circumstances the packet would just be routed to that connection. + // The only time this collision will occur if we receive the two Initial packets at the same time. + if added := s.tr.AddWithConnID(hdr.DestConnectionID, connID, conn); !added { + delete(s.zeroRTTQueues, hdr.DestConnectionID) + conn.closeWithTransportError(ConnectionRefused) + return nil + } + // Pass queued 0-RTT to the newly established connection. + if q, ok := s.zeroRTTQueues[hdr.DestConnectionID]; ok { + for _, p := range q.packets { + conn.handlePacket(p) + } + delete(s.zeroRTTQueues, hdr.DestConnectionID) + } + + s.handshakingCount.Go(func() { s.handleNewConn(conn) }) + go conn.run() + return nil +} + +func (s *baseServer) refuseNewConn(p receivedPacket, hdr *wire.Header) { + delete(s.zeroRTTQueues, hdr.DestConnectionID) + select { + case s.connectionRefusedQueue <- rejectedPacket{receivedPacket: p, hdr: hdr}: + default: + // drop packet if we can't send out the CONNECTION_REFUSED fast enough + p.buffer.Release() + } +} + +func (s *baseServer) handleNewConn(conn *wrappedConn) { + if s.acceptEarlyConns { + // wait until the early connection is ready, the handshake fails, or the server is closed + select { + case <-s.errorChan: + conn.closeWithTransportError(ConnectionRefused) + return + case <-conn.Context().Done(): + return + case <-conn.earlyConnReady(): + } + } else { + // wait until the handshake completes, fails, or the server is closed + select { + case <-s.errorChan: + conn.closeWithTransportError(ConnectionRefused) + return + case <-conn.Context().Done(): + return + case <-conn.HandshakeComplete(): + } + } + + select { + case s.connQueue <- conn.Conn: + default: + conn.closeWithTransportError(ConnectionRefused) + } +} + +func (s *baseServer) sendRetry(p rejectedPacket) { + if err := s.sendRetryPacket(p); err != nil { + s.logger.Debugf("Error sending Retry packet: %s", err) + } +} + +func (s *baseServer) sendRetryPacket(p rejectedPacket) error { + hdr := p.hdr + // Log the Initial packet now. + // If no Retry is sent, the packet will be logged by the connection. + (&wire.ExtendedHeader{Header: *hdr}).Log(s.logger) + srcConnID, err := s.connIDGenerator.GenerateConnectionID() + if err != nil { + return err + } + token, err := s.tokenGenerator.NewRetryToken(p.remoteAddr, hdr.DestConnectionID, srcConnID) + if err != nil { + return err + } + replyHdr := &wire.ExtendedHeader{} + replyHdr.Type = protocol.PacketTypeRetry + replyHdr.Version = hdr.Version + replyHdr.SrcConnectionID = srcConnID + replyHdr.DestConnectionID = hdr.SrcConnectionID + replyHdr.Token = token + if s.logger.Debug() { + s.logger.Debugf("Changing connection ID to %s.", srcConnID) + s.logger.Debugf("-> Sending Retry") + replyHdr.Log(s.logger) + } + + buf := getPacketBuffer() + defer buf.Release() + buf.Data, err = replyHdr.Append(buf.Data, hdr.Version) + if err != nil { + return err + } + // append the Retry integrity tag + tag := handshake.GetRetryIntegrityTag(buf.Data, hdr.DestConnectionID, hdr.Version) + buf.Data = append(buf.Data, tag[:]...) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketSent{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + SrcConnectionID: replyHdr.SrcConnectionID, + DestConnectionID: replyHdr.DestConnectionID, + Version: replyHdr.Version, + Token: &qlog.Token{Raw: token}, + }, + Raw: qlog.RawInfo{ + Length: len(buf.Data), + PayloadLength: int(replyHdr.Length), + }, + }) + } + _, err = s.conn.WritePacket(buf.Data, p.remoteAddr, p.info.OOB(), 0, protocol.ECNUnsupported) + return err +} + +func (s *baseServer) maybeSendInvalidToken(p rejectedPacket) { + defer p.buffer.Release() + + // Only send INVALID_TOKEN if we can unprotect the packet. + // This makes sure that we won't send it for packets that were corrupted. + hdr := p.hdr + sealer, opener := handshake.NewInitialAEAD(hdr.DestConnectionID, protocol.PerspectiveServer, hdr.Version) + data := p.data[:hdr.ParsedLen()+hdr.Length] + extHdr, err := unpackLongHeader(opener, hdr, data) + // Only send INVALID_TOKEN if we can unprotect the packet. + // This makes sure that we won't send it for packets that were corrupted. + if err != nil { + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + return + } + hdrLen := extHdr.ParsedLen() + if _, err := opener.Open(data[hdrLen:hdrLen], data[hdrLen:], extHdr.PacketNumber, data[:hdrLen]); err != nil { + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropPayloadDecryptError, + }) + } + return + } + if s.logger.Debug() { + s.logger.Debugf("Client sent an invalid retry token. Sending INVALID_TOKEN to %s.", p.remoteAddr) + } + if err := s.sendError(p.remoteAddr, hdr, sealer, InvalidToken, p.info); err != nil { + s.logger.Debugf("Error sending INVALID_TOKEN error: %s", err) + } +} + +func (s *baseServer) sendConnectionRefused(p rejectedPacket) { + defer p.buffer.Release() + sealer, _ := handshake.NewInitialAEAD(p.hdr.DestConnectionID, protocol.PerspectiveServer, p.hdr.Version) + if err := s.sendError(p.remoteAddr, p.hdr, sealer, ConnectionRefused, p.info); err != nil { + s.logger.Debugf("Error sending CONNECTION_REFUSED error: %s", err) + } +} + +// sendError sends the error as a response to the packet received with header hdr +func (s *baseServer) sendError(remoteAddr net.Addr, hdr *wire.Header, sealer handshake.LongHeaderSealer, errorCode qerr.TransportErrorCode, info packetInfo) error { + b := getPacketBuffer() + defer b.Release() + + ccf := &wire.ConnectionCloseFrame{ErrorCode: uint64(errorCode)} + + replyHdr := &wire.ExtendedHeader{} + replyHdr.Type = protocol.PacketTypeInitial + replyHdr.Version = hdr.Version + replyHdr.SrcConnectionID = hdr.DestConnectionID + replyHdr.DestConnectionID = hdr.SrcConnectionID + replyHdr.PacketNumberLen = protocol.PacketNumberLen4 + replyHdr.Length = 4 /* packet number len */ + ccf.Length(hdr.Version) + protocol.ByteCount(sealer.Overhead()) + var err error + b.Data, err = replyHdr.Append(b.Data, hdr.Version) + if err != nil { + return err + } + payloadOffset := len(b.Data) + + b.Data, err = ccf.Append(b.Data, hdr.Version) + if err != nil { + return err + } + + _ = sealer.Seal(b.Data[payloadOffset:payloadOffset], b.Data[payloadOffset:], replyHdr.PacketNumber, b.Data[:payloadOffset]) + b.Data = b.Data[0 : len(b.Data)+sealer.Overhead()] + + pnOffset := payloadOffset - int(replyHdr.PacketNumberLen) + sealer.EncryptHeader( + b.Data[pnOffset+4:pnOffset+4+16], + &b.Data[0], + b.Data[pnOffset:payloadOffset], + ) + + replyHdr.Log(s.logger) + wire.LogFrame(s.logger, ccf, true) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketSent{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + SrcConnectionID: replyHdr.SrcConnectionID, + DestConnectionID: replyHdr.DestConnectionID, + PacketNumber: replyHdr.PacketNumber, + Version: replyHdr.Version, + }, + Raw: qlog.RawInfo{ + Length: len(b.Data), + PayloadLength: int(replyHdr.Length), + }, + Frames: []qlog.Frame{{Frame: ccf}}, + }) + } + _, err = s.conn.WritePacket(b.Data, remoteAddr, info.OOB(), 0, protocol.ECNUnsupported) + return err +} + +func (s *baseServer) enqueueVersionNegotiationPacket(p receivedPacket) (bufferInUse bool) { + select { + case s.versionNegotiationQueue <- p: + return true + default: + // it's fine to not send version negotiation packets when we are busy + } + return false +} + +func (s *baseServer) maybeSendVersionNegotiationPacket(p receivedPacket) { + defer p.buffer.Release() + + v, err := wire.ParseVersion(p.data) + if err != nil { + s.logger.Debugf("failed to parse version for sending version negotiation packet: %s", err) + return + } + + _, src, dest, err := wire.ParseArbitraryLenConnectionIDs(p.data) + if err != nil { // should never happen + s.logger.Debugf("Dropping a packet with an unknown version for which we failed to parse connection IDs") + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }) + } + return + } + + s.logger.Debugf("Client offered version %s, sending Version Negotiation", v) + + data := wire.ComposeVersionNegotiation(dest, src, s.config.Versions) + if s.qlogger != nil { + s.qlogger.RecordEvent(qlog.VersionNegotiationSent{ + Header: qlog.PacketHeaderVersionNegotiation{ + SrcConnectionID: src, + DestConnectionID: dest, + }, + SupportedVersions: s.config.Versions, + }) + } + if _, err := s.conn.WritePacket(data, p.remoteAddr, p.info.OOB(), 0, protocol.ECNUnsupported); err != nil { + s.logger.Debugf("Error sending Version Negotiation: %s", err) + } +} diff --git a/third_party/quic-go/server_test.go b/third_party/quic-go/server_test.go new file mode 100644 index 0000000..4d6db4d --- /dev/null +++ b/third_party/quic-go/server_test.go @@ -0,0 +1,1420 @@ +package quic + +import ( + "context" + "crypto/rand" + "crypto/tls" + "errors" + "net" + "slices" + "testing" + "time" + + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type testServer struct{ *baseServer } + +type serverOpts struct { + eventRecorder *events.Recorder + config *Config + tokenGeneratorKey TokenGeneratorKey + maxTokenAge time.Duration + useRetry bool + disableVersionNegotiation bool + acceptEarly bool + newConn func( + context.Context, + context.CancelCauseFunc, + sendConn, + connRunner, + protocol.ConnectionID, // original dest connection ID + *protocol.ConnectionID, // retry src connection ID + protocol.ConnectionID, // client dest connection ID + protocol.ConnectionID, // destination connection ID + protocol.ConnectionID, // source connection ID + ConnectionIDGenerator, + *statelessResetter, + *Config, + *tls.Config, + *handshake.TokenGenerator, + bool, /* client address validated by an address validation token */ + time.Duration, + qlogwriter.Trace, + utils.Logger, + protocol.Version, + ) *wrappedConn +} + +func newTestServer(t *testing.T, serverOpts *serverOpts) *testServer { + t.Helper() + c, err := wrapConn(newUDPConnLocalhost(t), false) + require.NoError(t, err) + verifySourceAddress := func(net.Addr) bool { return serverOpts.useRetry } + config := populateConfig(serverOpts.config) + tr := &Transport{Conn: newUDPConnLocalhost(t)} + tr.init(true) + s := newServer( + c, + (*packetHandlerMap)(tr), + &protocol.DefaultConnectionIDGenerator{}, + &statelessResetter{}, + func(ctx context.Context, _ *ClientInfo) (context.Context, error) { return ctx, nil }, + &tls.Config{}, + config, + serverOpts.eventRecorder, + func() {}, + serverOpts.tokenGeneratorKey, + serverOpts.maxTokenAge, + verifySourceAddress, + serverOpts.disableVersionNegotiation, + serverOpts.acceptEarly, + ) + s.newConn = serverOpts.newConn + t.Cleanup(func() { s.Close() }) + return &testServer{s} +} + +func getLongHeaderPacketEncrypted(t *testing.T, remoteAddr net.Addr, extHdr *wire.ExtendedHeader, data []byte) receivedPacket { + t.Helper() + hdr := extHdr.Header + if hdr.Type != protocol.PacketTypeInitial { + t.Fatal("can only encrypt Initial packets") + } + p := getLongHeaderPacket(t, remoteAddr, extHdr, data) + sealer, _ := handshake.NewInitialAEAD(hdr.DestConnectionID, protocol.PerspectiveClient, hdr.Version) + n := len(p.data) - len(data) // length of the header + p.data = slices.Grow(p.data, 16) + _ = sealer.Seal(p.data[n:n], p.data[n:], extHdr.PacketNumber, p.data[:n]) + p.data = p.data[:len(p.data)+16] + sealer.EncryptHeader(p.data[n:n+16], &p.data[0], p.data[n-int(extHdr.PacketNumberLen):n]) + return p +} + +func randConnID(l int) protocol.ConnectionID { + b := make([]byte, l) + rand.Read(b) + return protocol.ParseConnectionID(b) +} + +func getValidInitialPacket(t *testing.T, raddr net.Addr, srcConnID, destConnID protocol.ConnectionID) receivedPacket { + t.Helper() + return getLongHeaderPacket(t, + raddr, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + SrcConnectionID: srcConnID, + DestConnectionID: destConnID, + Length: protocol.MinInitialPacketSize, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, protocol.MinInitialPacketSize), + ) +} + +// checkConnectionClose checks +// 1. the arguments of the SentPacket tracer call, and +// 2. reads and parses the packet sent by the server +func checkConnectionClose( + t *testing.T, + conn *net.UDPConn, + eventRecorder *events.Recorder, + expectedSrcConnID protocol.ConnectionID, + expectedDestConnID protocol.ConnectionID, + expectedErrorCode qerr.TransportErrorCode, +) { + t.Helper() + + conn.SetReadDeadline(time.Now().Add(time.Second)) + b := make([]byte, 1500) + n, _, err := conn.ReadFromUDP(b) + require.NoError(t, err) + parsedHdr, _, _, err := wire.ParsePacket(b[:n]) + require.NoError(t, err) + require.Equal(t, protocol.PacketTypeInitial, parsedHdr.Type) + require.Equal(t, expectedSrcConnID, parsedHdr.SrcConnectionID) + require.Equal(t, expectedDestConnID, parsedHdr.DestConnectionID) + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketSent{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + SrcConnectionID: expectedSrcConnID, + DestConnectionID: expectedDestConnID, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: n, PayloadLength: int(parsedHdr.Length)}, + Frames: []qlog.Frame{ + {Frame: &qlog.ConnectionCloseFrame{ErrorCode: uint64(expectedErrorCode)}}, + }, + }, + }, + eventRecorder.Events(qlog.PacketSent{}), + ) +} + +func checkRetry(t *testing.T, + conn *net.UDPConn, + eventRecorder *events.Recorder, + expectedDestConnID protocol.ConnectionID, +) { + t.Helper() + + conn.SetReadDeadline(time.Now().Add(time.Second)) + b := make([]byte, 1500) + n, _, err := conn.ReadFromUDP(b) + require.NoError(t, err) + parsedHdr, _, _, err := wire.ParsePacket(b[:n]) + require.NoError(t, err) + require.Equal(t, protocol.PacketTypeRetry, parsedHdr.Type) + require.Equal(t, expectedDestConnID, parsedHdr.DestConnectionID) + require.NotNil(t, parsedHdr.Token) + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketSent{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeRetry, + DestConnectionID: expectedDestConnID, + SrcConnectionID: parsedHdr.SrcConnectionID, + Version: parsedHdr.Version, + Token: &qlog.Token{Raw: parsedHdr.Token}, + }, + Raw: qlog.RawInfo{Length: n}, + }, + }, + eventRecorder.Events(qlog.PacketSent{}), + ) +} + +func TestListen(t *testing.T) { + _, err := ListenAddr("localhost:0", nil, nil) + require.Error(t, err) + require.Contains(t, err.Error(), "quic: tls.Config not set") + + _, err = Listen(nil, &tls.Config{}, &Config{Versions: []protocol.Version{0x1234}}) + require.Error(t, err) + require.Contains(t, err.Error(), "invalid QUIC version: 0x1234") +} + +func TestListenAddr(t *testing.T) { + _, err := ListenAddr("127.0.0.1", &tls.Config{}, &Config{}) + require.Error(t, err) + require.IsType(t, &net.AddrError{}, err) + + _, err = ListenAddr("1.1.1.1:1111", &tls.Config{}, &Config{}) + require.Error(t, err) + require.IsType(t, &net.OpError{}, err) + + ln, err := ListenAddr("127.0.0.1:0", &tls.Config{}, &Config{}) + require.NoError(t, err) + defer ln.Close() +} + +func TestServerPacketDropping(t *testing.T) { + t.Run("destination connection ID too short", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + testServerDroppedPacket(t, + conn, + getValidInitialPacket(t, conn.LocalAddr(), randConnID(5), randConnID(7)), + protocol.Version1, + qlog.PacketTypeInitial, + qlog.PacketDropUnexpectedPacket, + ) + }) + + t.Run("Initial packet too small", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + p := getLongHeaderPacket(t, + conn.LocalAddr(), + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: randConnID(8), + Version: protocol.Version1, + }, + PacketNumberLen: 2, + }, + make([]byte, protocol.MinInitialPacketSize-100), + ) + require.Greater(t, len(p.data), protocol.MinInitialPacketSize-100) + require.Less(t, len(p.data), protocol.MinInitialPacketSize) + testServerDroppedPacket(t, + conn, + p, + protocol.Version1, + qlog.PacketTypeInitial, + qlog.PacketDropUnexpectedPacket, + ) + }) + + // we should not send a Version Negotiation packet if the packet is smaller than 1200 bytes + t.Run("packet of unknown version, too small", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + p := getLongHeaderPacket(t, + conn.LocalAddr(), + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + DestConnectionID: randConnID(8), + Version: 0x42, + }, + PacketNumberLen: 2, + }, + make([]byte, protocol.MinUnknownVersionPacketSize-100), + ) + require.Greater(t, len(p.data), protocol.MinUnknownVersionPacketSize-100) + require.Less(t, len(p.data), protocol.MinUnknownVersionPacketSize) + testServerDroppedPacket(t, + conn, + p, + 0x42, + "", + qlog.PacketDropUnexpectedPacket, + ) + }) + + t.Run("not an Initial packet", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + testServerDroppedPacket(t, + conn, + getLongHeaderPacket(t, + conn.LocalAddr(), + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Version: protocol.Version1, + }, + PacketNumberLen: 2, + }, + nil, + ), + protocol.Version1, + qlog.PacketTypeHandshake, + qlog.PacketDropUnexpectedPacket, + ) + }) + + // as a server, we should never receive a Version Negotiation packet + t.Run("Version Negotiation packet", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + data := wire.ComposeVersionNegotiation( + protocol.ArbitraryLenConnectionID{1, 2, 3, 4}, + protocol.ArbitraryLenConnectionID{4, 3, 2, 1}, + []protocol.Version{1, 2, 3}, + ) + testServerDroppedPacket(t, + conn, + receivedPacket{ + remoteAddr: conn.LocalAddr(), + data: data, + buffer: getPacketBuffer(), + }, + 0, // version negotiation packets don't have a version + qlog.PacketTypeVersionNegotiation, + qlog.PacketDropUnexpectedPacket, + ) + }) +} + +func testServerDroppedPacket(t *testing.T, + conn *net.UDPConn, + p receivedPacket, + expectedVersion qlog.Version, + expectedPacketType qlog.PacketType, + expectedDropReason qlog.PacketDropReason, +) { + readChan := make(chan struct{}) + go func() { + defer close(readChan) + conn.ReadFrom(make([]byte, 1000)) + }() + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{eventRecorder: &eventRecorder}) + + server.handlePacket(p) + + select { + case <-readChan: + t.Fatal("didn't expect to receive a packet") + case <-time.After(scaleDuration(5 * time.Millisecond)): + } + + var expectedPacketNumber protocol.PacketNumber + if expectedPacketType != qlog.PacketTypeVersionNegotiation && expectedPacketType != "" { + expectedPacketNumber = protocol.InvalidPacketNumber + } + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: expectedPacketType, + PacketNumber: expectedPacketNumber, + Version: expectedVersion, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: expectedDropReason, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) +} + +func TestServerVersionNegotiation(t *testing.T) { + t.Run("enabled", func(t *testing.T) { + testServerVersionNegotiation(t, true) + }) + t.Run("disabled", func(t *testing.T) { + testServerVersionNegotiation(t, false) + }) +} + +func testServerVersionNegotiation(t *testing.T, enabled bool) { + conn := newUDPConnLocalhost(t) + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + eventRecorder: &eventRecorder, + disableVersionNegotiation: !enabled, + }) + + srcConnID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5}) + destConnID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6}) + packet := getLongHeaderPacket(t, conn.LocalAddr(), + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + SrcConnectionID: srcConnID, + DestConnectionID: destConnID, + Version: 0x42, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, protocol.MinUnknownVersionPacketSize), + ) + + written := make(chan []byte, 1) + go func() { + b := make([]byte, 1500) + n, _, _ := conn.ReadFrom(b) + written <- b[:n] + }() + server.handlePacket(packet) + + switch enabled { + case true: + select { + case b := <-written: + require.True(t, wire.IsVersionNegotiationPacket(b)) + dest, src, versions, err := wire.ParseVersionNegotiationPacket(b) + require.NoError(t, err) + require.Equal(t, protocol.ArbitraryLenConnectionID(srcConnID.Bytes()), dest) + require.Equal(t, protocol.ArbitraryLenConnectionID(destConnID.Bytes()), src) + require.NotContains(t, versions, protocol.Version(0x42)) + + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionNegotiationSent{ + Header: qlog.PacketHeaderVersionNegotiation{ + SrcConnectionID: src, + DestConnectionID: dest, + }, + SupportedVersions: server.config.Versions, + }, + }, + eventRecorder.Events(qlog.VersionNegotiationSent{}), + ) + case <-time.After(time.Second): + t.Fatal("timeout") + } + case false: + select { + case <-written: + t.Fatal("expected no version negotiation packet") + case <-time.After(scaleDuration(10 * time.Millisecond)): + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{Version: 0x42}, + Raw: qlog.RawInfo{Length: int(packet.Size())}, + Trigger: qlog.PacketDropUnexpectedVersion, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + } + } +} + +func TestServerRetry(t *testing.T) { + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{eventRecorder: &eventRecorder, useRetry: true}) + conn := newUDPConnLocalhost(t) + + packet := getLongHeaderPacket(t, conn.LocalAddr(), + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + SrcConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, protocol.MinUnknownVersionPacketSize), + ) + + server.handlePacket(packet) + checkRetry(t, conn, &eventRecorder, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5})) +} + +func TestServerTokenValidation(t *testing.T) { + var tokenGeneratorKey handshake.TokenProtectorKey + rand.Read(tokenGeneratorKey[:]) + tg := handshake.NewTokenGenerator(tokenGeneratorKey) + + t.Run("retry token with invalid address", func(t *testing.T) { + token, err := tg.NewRetryToken( + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 1337}, + protocol.ConnectionID{}, + protocol.ConnectionID{}, + ) + require.NoError(t, err) + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + useRetry: true, + eventRecorder: &eventRecorder, + tokenGeneratorKey: tokenGeneratorKey, + }) + + testServerTokenValidation(t, server, &eventRecorder, newUDPConnLocalhost(t), token, false, true, false) + }) + + t.Run("expired retry token", func(t *testing.T) { + conn := newUDPConnLocalhost(t) + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + useRetry: true, + eventRecorder: &eventRecorder, + config: &Config{HandshakeIdleTimeout: time.Millisecond / 2}, + tokenGeneratorKey: tokenGeneratorKey, + }) + + token, err := tg.NewRetryToken(conn.LocalAddr(), protocol.ConnectionID{}, protocol.ConnectionID{}) + require.NoError(t, err) + // the maximum retry token age is equivalent to the handshake timeout + time.Sleep(time.Millisecond) // make sure the token is expired + testServerTokenValidation(t, server, &eventRecorder, conn, token, false, true, false) + }) + + // if the packet is corrupted, it will just be dropped (no INVALID_TOKEN nor Retry is sent) + t.Run("corrupted packet", func(t *testing.T) { + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + useRetry: true, + eventRecorder: &eventRecorder, + config: &Config{HandshakeIdleTimeout: time.Millisecond / 2}, + tokenGeneratorKey: tokenGeneratorKey, + }) + + conn := newUDPConnLocalhost(t) + token, err := tg.NewRetryToken(conn.LocalAddr(), protocol.ConnectionID{}, protocol.ConnectionID{}) + require.NoError(t, err) + time.Sleep(time.Millisecond) // make sure the token is expired + testServerTokenValidation(t, server, &eventRecorder, conn, token, true, false, true) + }) + + t.Run("invalid non-retry token", func(t *testing.T) { + var tokenGeneratorKey2 handshake.TokenProtectorKey + rand.Read(tokenGeneratorKey2[:]) + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + tokenGeneratorKey: tokenGeneratorKey2, // use a different key + useRetry: true, + eventRecorder: &eventRecorder, + maxTokenAge: time.Millisecond, + }) + + conn := newUDPConnLocalhost(t) + token, err := tg.NewToken(conn.LocalAddr(), 10*time.Millisecond) + require.NoError(t, err) + time.Sleep(3 * time.Millisecond) // make sure the token is expired + testServerTokenValidation(t, server, &eventRecorder, conn, token, false, false, true) + }) + + t.Run("expired non-retry token", func(t *testing.T) { + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + tokenGeneratorKey: tokenGeneratorKey, + useRetry: true, + eventRecorder: &eventRecorder, + maxTokenAge: time.Millisecond, + }) + + conn := newUDPConnLocalhost(t) + token, err := tg.NewToken(conn.LocalAddr(), 100*time.Millisecond) + require.NoError(t, err) + time.Sleep(3 * time.Millisecond) // make sure the token is expired + testServerTokenValidation(t, server, &eventRecorder, conn, token, false, false, true) + }) +} + +func testServerTokenValidation( + t *testing.T, + server *testServer, + eventRecorder *events.Recorder, + conn *net.UDPConn, + token []byte, + corruptedPacket bool, + expectInvalidTokenConnectionClose bool, + expectRetry bool, +) { + hdr := wire.Header{ + Type: protocol.PacketTypeInitial, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + Token: token, + Length: protocol.MinInitialPacketSize + protocol.ByteCount(protocol.PacketNumberLen4) + 16, + Version: protocol.Version1, + } + packet := getLongHeaderPacketEncrypted(t, + conn.LocalAddr(), + &wire.ExtendedHeader{Header: hdr, PacketNumberLen: protocol.PacketNumberLen4}, + make([]byte, protocol.MinInitialPacketSize), + ) + if corruptedPacket { + packet.data[len(packet.data)-10] ^= 0xff // corrupt the packet + server.handlePacket(packet) + + require.Eventually(t, + func() bool { return len(eventRecorder.Events(qlog.PacketDropped{})) > 0 }, + time.Second, + 10*time.Millisecond, + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeInitial, + PacketNumber: protocol.InvalidPacketNumber, + Version: hdr.Version, + }, + Raw: qlog.RawInfo{Length: int(packet.Size())}, + Trigger: qlog.PacketDropPayloadDecryptError, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + return + } + + server.handlePacket(packet) + + if expectInvalidTokenConnectionClose { + checkConnectionClose(t, conn, eventRecorder, hdr.DestConnectionID, hdr.SrcConnectionID, qerr.InvalidToken) + } + if expectRetry { + checkRetry(t, conn, eventRecorder, hdr.SrcConnectionID) + } +} + +type connConstructorArgs struct { + ctx context.Context + connRunner connRunner + config *Config + origDestConnID protocol.ConnectionID + retrySrcConnID *protocol.ConnectionID + clientDestConnID protocol.ConnectionID + destConnID protocol.ConnectionID + srcConnID protocol.ConnectionID +} + +type connConstructorRecorder struct { + ch chan connConstructorArgs + + hooks []*connTestHooks +} + +func newConnConstructorRecorder(hooks ...*connTestHooks) *connConstructorRecorder { + return &connConstructorRecorder{ + ch: make(chan connConstructorArgs, len(hooks)), + hooks: hooks, + } +} + +func (r *connConstructorRecorder) Args() <-chan connConstructorArgs { return r.ch } + +func (r *connConstructorRecorder) NewConn( + ctx context.Context, + _ context.CancelCauseFunc, + _ sendConn, + connRunner connRunner, + origDestConnID protocol.ConnectionID, + retrySrcConnID *protocol.ConnectionID, + clientDestConnID protocol.ConnectionID, + destConnID protocol.ConnectionID, + srcConnID protocol.ConnectionID, + _ ConnectionIDGenerator, + _ *statelessResetter, + config *Config, + _ *tls.Config, + _ *handshake.TokenGenerator, + _ bool, + _ time.Duration, + _ qlogwriter.Trace, + _ utils.Logger, + _ protocol.Version, +) *wrappedConn { + r.ch <- connConstructorArgs{ + ctx: ctx, + connRunner: connRunner, + config: config, + origDestConnID: origDestConnID, + retrySrcConnID: retrySrcConnID, + clientDestConnID: clientDestConnID, + destConnID: destConnID, + srcConnID: srcConnID, + } + hooks := r.hooks[0] + r.hooks = r.hooks[1:] + return &wrappedConn{testHooks: hooks} +} + +func TestServerCreateConnection(t *testing.T) { + t.Run("without retry", func(t *testing.T) { + testServerCreateConnection(t, false) + }) + t.Run("with retry", func(t *testing.T) { + testServerCreateConnection(t, true) + }) +} + +func testServerCreateConnection(t *testing.T, useRetry bool) { + tokenGeneratorKey := TokenGeneratorKey{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11} + tg := handshake.NewTokenGenerator(tokenGeneratorKey) + + server := newTestServer(t, &serverOpts{ + useRetry: useRetry, + tokenGeneratorKey: tokenGeneratorKey, + }) + + done := make(chan struct{}, 3) + handledPackets := make(chan receivedPacket, 1) + recorder := newConnConstructorRecorder(&connTestHooks{ + run: func() error { done <- struct{}{}; return nil }, + context: func() context.Context { done <- struct{}{}; return context.Background() }, + handshakeComplete: func() <-chan struct{} { done <- struct{}{}; return make(chan struct{}) }, + handlePacket: func(p receivedPacket) { handledPackets <- p }, + }) + server.newConn = recorder.NewConn + + conn := newUDPConnLocalhost(t) + var token []byte + if useRetry { + var err error + token, err = tg.NewRetryToken( + conn.LocalAddr(), + protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}), + protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), + ) + require.NoError(t, err) + } + hdr := wire.Header{ + Type: protocol.PacketTypeInitial, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), + Length: protocol.MinInitialPacketSize + protocol.ByteCount(protocol.PacketNumberLen4) + 16, + Token: token, + Version: protocol.Version1, + } + packet := getLongHeaderPacketEncrypted(t, + conn.LocalAddr(), + &wire.ExtendedHeader{Header: hdr, PacketNumberLen: protocol.PacketNumberLen4}, + make([]byte, protocol.MinInitialPacketSize), + ) + + server.handlePacket(packet) + + select { + case p := <-handledPackets: + require.Equal(t, packet, p) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + var args connConstructorArgs + select { + case args = <-recorder.Args(): + case <-time.After(time.Second): + t.Fatal("timeout") + } + + assert.Equal(t, protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), args.destConnID) + assert.NotEqual(t, args.origDestConnID, args.srcConnID) + if useRetry { + assert.Equal(t, protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), args.destConnID) + assert.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xad, 0xc0, 0xde}), args.origDestConnID) + assert.Equal(t, protocol.ParseConnectionID([]byte{0xde, 0xca, 0xfb, 0xad}), *args.retrySrcConnID) + } else { + assert.Equal(t, protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}), args.origDestConnID) + assert.Zero(t, args.retrySrcConnID) + } + + for range 3 { + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + } +} + +func TestServerClose(t *testing.T) { + var hooks []*connTestHooks + const numConns = 3 + done := make(chan struct{}, numConns) + for range numConns { + hooks = append(hooks, &connTestHooks{ + closeWithTransportError: func(TransportErrorCode) { done <- struct{}{} }, + }) + } + recorder := newConnConstructorRecorder(hooks...) + server := newTestServer(t, &serverOpts{newConn: recorder.NewConn}) + + for range numConns { + b := make([]byte, 10) + rand.Read(b) + connID := protocol.ParseConnectionID(b) + server.handlePacket(getValidInitialPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + randConnID(6), + connID, + )) + select { + case <-recorder.Args(): + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + + server.Close() + // closing closes all handshaking connections with CONNECTION_REFUSED + for range numConns { + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + + // Accept returns ErrServerClosed after closing + for range 5 { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err := server.Accept(ctx) + require.ErrorIs(t, err, ErrServerClosed) + require.ErrorIs(t, err, net.ErrClosed) + } +} + +func TestServerGetConfigForClientAccept(t *testing.T) { + recorder := newConnConstructorRecorder(&connTestHooks{}) + server := newTestServer(t, &serverOpts{ + config: &Config{ + GetConfigForClient: func(*ClientInfo) (*Config, error) { + return &Config{MaxIncomingStreams: 1234}, nil + }, + }, + newConn: recorder.NewConn, + }) + + conn := newUDPConnLocalhost(t) + packet := getValidInitialPacket(t, + conn.LocalAddr(), + protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + ) + + server.handlePacket(packet) + + var args connConstructorArgs + select { + case args = <-recorder.Args(): + require.EqualValues(t, 1234, args.config.MaxIncomingStreams) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + assert.Equal(t, protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), args.destConnID) + assert.NotEqual(t, args.origDestConnID, args.srcConnID) +} + +func TestServerGetConfigForClientReject(t *testing.T) { + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + eventRecorder: &eventRecorder, + config: &Config{ + GetConfigForClient: func(*ClientInfo) (*Config, error) { + return nil, errors.New("rejected") + }, + }, + }) + + conn := newUDPConnLocalhost(t) + srcConnID := randConnID(6) + destConnID := randConnID(8) + server.handlePacket(getValidInitialPacket(t, conn.LocalAddr(), srcConnID, destConnID)) + + checkConnectionClose(t, conn, &eventRecorder, destConnID, srcConnID, qerr.ConnectionRefused) +} + +type failingConnectionIDGenerator struct{ err error } + +func (g failingConnectionIDGenerator) GenerateConnectionID() (ConnectionID, error) { + return ConnectionID{}, g.err +} + +func (failingConnectionIDGenerator) ConnectionIDLen() int { return 8 } + +func TestServerCancelsConnContextWhenConnectionIDGenerationFails(t *testing.T) { + want := errors.New("connection ID randomness failed") + server := newTestServer(t, &serverOpts{}) + server.connIDGenerator = failingConnectionIDGenerator{err: want} + canceled := make(chan error, 1) + server.connContext = func(ctx context.Context, _ *ClientInfo) (context.Context, error) { + derived, cancel := context.WithCancelCause(ctx) + go func() { + <-derived.Done() + canceled <- context.Cause(derived) + }() + _ = cancel // cancellation is owned by the server's derived context + return derived, nil + } + + conn := newUDPConnLocalhost(t) + server.handlePacket(getValidInitialPacket(t, conn.LocalAddr(), randConnID(6), randConnID(8))) + select { + case cause := <-canceled: + require.ErrorIs(t, cause, want) + case <-time.After(time.Second): + t.Fatal("ConnContext was not canceled after connection ID generation failed") + } +} + +func TestServerReceiveQueue(t *testing.T) { + var eventRecorder events.Recorder + acceptConn := make(chan struct{}) + defer close(acceptConn) + newConnChan := make(chan struct{}, protocol.MaxServerUnprocessedPackets+2) + server := newTestServer(t, &serverOpts{ + eventRecorder: &eventRecorder, + newConn: func( + _ context.Context, + _ context.CancelCauseFunc, + _ sendConn, + _ connRunner, + _ protocol.ConnectionID, + _ *protocol.ConnectionID, + _ protocol.ConnectionID, + _ protocol.ConnectionID, + _ protocol.ConnectionID, + _ ConnectionIDGenerator, + _ *statelessResetter, + _ *Config, + _ *tls.Config, + _ *handshake.TokenGenerator, + _ bool, + _ time.Duration, + _ qlogwriter.Trace, + _ utils.Logger, + _ protocol.Version, + ) *wrappedConn { + newConnChan <- struct{}{} + <-acceptConn + return &wrappedConn{testHooks: &connTestHooks{handlePacket: func(receivedPacket) {}}} + }, + }) + + conn := newUDPConnLocalhost(t) + for i := range protocol.MaxServerUnprocessedPackets + 1 { + server.handlePacket(getValidInitialPacket(t, conn.LocalAddr(), randConnID(6), randConnID(8))) + // newConn blocks on the acceptConn channel, so this blocks the server's run loop + if i == 0 { + select { + case <-newConnChan: + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + } + + p := getValidInitialPacket(t, conn.LocalAddr(), randConnID(6), randConnID(8)) + server.handlePacket(p) + + require.Eventually(t, + func() bool { return len(eventRecorder.Events(qlog.PacketDropped{})) > 0 }, + time.Second, + 10*time.Millisecond, + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) +} + +func TestServerAccept(t *testing.T) { + t.Run("without accept early", func(t *testing.T) { + testServerAccept(t, false) + }) + t.Run("with accept early", func(t *testing.T) { + testServerAccept(t, true) + }) +} + +func testServerAccept(t *testing.T, acceptEarly bool) { + ready := make(chan struct{}) + hooks := &connTestHooks{} + if acceptEarly { + hooks.earlyConnReady = func() <-chan struct{} { return ready } + } else { + hooks.handshakeComplete = func() <-chan struct{} { return ready } + } + recorder := newConnConstructorRecorder(hooks) + server := newTestServer(t, &serverOpts{ + acceptEarly: acceptEarly, + newConn: recorder.NewConn, + }) + + // Accept should respect the context + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, err := server.Accept(ctx) + require.ErrorIs(t, err, context.Canceled) + + // establish a new connection, which then starts handshaking + server.handlePacket(getValidInitialPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + randConnID(6), + randConnID(8), + )) + + accepted := make(chan error, 1) + go func() { + _, err := server.Accept(context.Background()) + accepted <- err + }() + + select { + case <-accepted: + t.Fatal("server accepted the connection too early") + case <-time.After(scaleDuration(5 * time.Millisecond)): + } + // now complete the handshake + close(ready) + + select { + case err := <-accepted: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestServerAcceptHandshakeFailure(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + recorder := newConnConstructorRecorder(&connTestHooks{ + context: func() context.Context { return ctx }, + handshakeComplete: func() <-chan struct{} { return make(chan struct{}) }, + }) + server := newTestServer(t, &serverOpts{newConn: recorder.NewConn}) + + // establish a new connection, which then starts handshaking + server.handlePacket(getValidInitialPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + randConnID(6), + randConnID(8), + )) + + accepted := make(chan error, 1) + go func() { + _, err := server.Accept(context.Background()) + accepted <- err + }() + + cancel() + select { + case <-accepted: + t.Fatal("server should not have accepted the connection") + case <-time.After(scaleDuration(5 * time.Millisecond)): + } +} + +func TestServerAcceptQueue(t *testing.T) { + var conns []*connTestHooks + rejectedCloseError := make(chan TransportErrorCode, 1) + for i := range protocol.MaxAcceptQueueSize + 2 { + conn := &connTestHooks{ + handshakeComplete: func() <-chan struct{} { + c := make(chan struct{}) + close(c) + return c + }, + } + conns = append(conns, conn) + if i == protocol.MaxAcceptQueueSize { + conn.closeWithTransportError = func(code TransportErrorCode) { rejectedCloseError <- code } + continue + } + } + recorder := newConnConstructorRecorder(conns...) + server := newTestServer(t, &serverOpts{newConn: recorder.NewConn}) + + for range protocol.MaxAcceptQueueSize { + b := make([]byte, 16) + rand.Read(b) + connID := protocol.ParseConnectionID(b) + server.handlePacket( + getValidInitialPacket(t, &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, randConnID(6), connID), + ) + select { + case args := <-recorder.Args(): + require.Equal(t, connID, args.origDestConnID) + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + // wait for the connection to be enqueued + time.Sleep(scaleDuration(10 * time.Millisecond)) + + server.handlePacket( + getValidInitialPacket(t, &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, randConnID(6), randConnID(8)), + ) + select { + case <-recorder.Args(): + case <-time.After(time.Second): + t.Fatal("timeout") + } + select { + case code := <-rejectedCloseError: + require.Equal(t, ConnectionRefused, code) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + // accept one connection, freeing up one slot in the accept queue + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err := server.Accept(ctx) + require.NoError(t, err) + + // it's now possible to enqueue a new connection + server.handlePacket( + getValidInitialPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + randConnID(6), + protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), + ), + ) + select { + case args := <-recorder.Args(): + require.Equal(t, protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), args.origDestConnID) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestServer0RTTReordering(t *testing.T) { + var eventRecorder events.Recorder + packets := make(chan receivedPacket, protocol.Max0RTTQueueLen+1) + done := make(chan struct{}) + recorder := newConnConstructorRecorder(&connTestHooks{ + handlePacket: func(p receivedPacket) { packets <- p }, + earlyConnReady: func() <-chan struct{} { return make(chan struct{}) }, + run: func() error { close(done); return nil }, + }) + server := newTestServer(t, &serverOpts{ + acceptEarly: true, + eventRecorder: &eventRecorder, + newConn: recorder.NewConn, + }) + + connID := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}) + + var zeroRTTPackets []receivedPacket + + for range protocol.Max0RTTQueueLen { + p := getLongHeaderPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketType0RTT, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: connID, + Length: 100, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, 100), + ) + server.handlePacket(p) + zeroRTTPackets = append(zeroRTTPackets, p) + } + + // send one more packet, this one should be dropped + p := getLongHeaderPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketType0RTT, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: connID, + Length: 100, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, 100), + ) + server.handlePacket(p) + + require.Eventually(t, + func() bool { return len(eventRecorder.Events(qlog.PacketDropped{})) > 0 }, + time.Second, + 10*time.Millisecond, + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + + // now receive the Initial + initial := getValidInitialPacket(t, &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, randConnID(5), connID) + server.handlePacket(initial) + + for i := range protocol.Max0RTTQueueLen + 1 { + select { + case p := <-packets: + if i == 0 { + require.Equal(t, initial.data, p.data) + } else { + require.Equal(t, zeroRTTPackets[i-1], p) + } + case <-time.After(time.Second): + t.Fatal("timeout") + } + } + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestServer0RTTQueueing(t *testing.T) { + var eventRecorder events.Recorder + server := newTestServer(t, &serverOpts{ + acceptEarly: true, + eventRecorder: &eventRecorder, + }) + + firstRcvTime := monotime.Now() + otherRcvTime := firstRcvTime.Add(protocol.Max0RTTQueueingDuration / 2) + var sizes []protocol.ByteCount + for i := range protocol.Max0RTTQueues { + b := make([]byte, 16) + rand.Read(b) + connID := protocol.ParseConnectionID(b) + size := protocol.ByteCount(500 + i) + p := getLongHeaderPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketType0RTT, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: connID, + Length: size, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, size), + ) + if i == 0 { + p.rcvTime = firstRcvTime + } else { + p.rcvTime = otherRcvTime + } + sizes = append(sizes, p.Size()) + server.handlePacket(p) + } + + // maximum number of 0-RTT queues is reached, further packets are dropped + p := getLongHeaderPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketType0RTT, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}), + Length: 123, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, 123), + ) + server.handlePacket(p) + require.Eventually(t, + func() bool { return len(eventRecorder.Events(qlog.PacketDropped{})) > 0 }, + time.Second, + 10*time.Millisecond, + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + // There's no cleanup Go routine. + // Cleanup is triggered when new packets are received. + // 1. Receive one handshake packet, which triggers the cleanup of the first 0-RTT queue + triggerPacket := getLongHeaderPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), + Length: 123, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, 123), + ) + triggerPacket.rcvTime = firstRcvTime.Add(protocol.Max0RTTQueueingDuration + time.Nanosecond) + server.handlePacket(triggerPacket) + require.Eventually(t, + func() bool { return len(eventRecorder.Events(qlog.PacketDropped{})) == 2 }, + time.Second, + 10*time.Millisecond, + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + PacketNumber: protocol.InvalidPacketNumber, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(triggerPacket.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(sizes[0])}, + Trigger: qlog.PacketDropDOSPrevention, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + // 2. Receive another handshake packet, which triggers the cleanup of the other 0-RTT queues + triggerPacket = getLongHeaderPacket(t, + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 42}, + &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeHandshake, + SrcConnectionID: protocol.ParseConnectionID([]byte{5, 4, 3, 2, 1}), + DestConnectionID: protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}), + Length: 124, + Version: protocol.Version1, + }, + PacketNumberLen: protocol.PacketNumberLen4, + }, + make([]byte, 124), + ) + triggerPacket.rcvTime = otherRcvTime.Add(protocol.Max0RTTQueueingDuration + time.Nanosecond) + server.handlePacket(triggerPacket) + + expectedEvents := []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketTypeHandshake, + PacketNumber: protocol.InvalidPacketNumber, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(triggerPacket.Size())}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + } + for i := range protocol.Max0RTTQueues - 1 { + expectedEvents = append(expectedEvents, qlog.PacketDropped{ + Header: qlog.PacketHeader{ + PacketType: qlog.PacketType0RTT, + PacketNumber: protocol.InvalidPacketNumber, + Version: protocol.Version1, + }, + Raw: qlog.RawInfo{Length: int(sizes[i+1])}, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + require.Eventually(t, + func() bool { return len(eventRecorder.Events(qlog.PacketDropped{})) == len(expectedEvents) }, + time.Second, + 10*time.Millisecond, + ) + + // queues are dropped in random order + for _, event := range expectedEvents { + require.Contains(t, eventRecorder.Events(qlog.PacketDropped{}), event) + } +} diff --git a/third_party/quic-go/sni.go b/third_party/quic-go/sni.go new file mode 100644 index 0000000..f63023f --- /dev/null +++ b/third_party/quic-go/sni.go @@ -0,0 +1,136 @@ +package quic + +import ( + "encoding/binary" + "errors" + "io" +) + +const ( + extTypeSNI = 0 + extTypeECH = 0xfe0d +) + +// findSNIAndECH parses the given byte slice as a ClientHello, and locates: +// - the position and length of the Server Name Indication (SNI) extension, +// - the position of the Encrypted Client Hello (ECH) extension. +// If no SNI extension is found, it returns -1 for the SNI position. +// If no ECH extension is found, it returns -1 for the ECH position. +func findSNIAndECH(data []byte) (sniPos, sniLen, echPos int, err error) { + if len(data) < 4 { + return 0, 0, 0, io.ErrUnexpectedEOF + } + if data[0] != 1 { + return 0, 0, 0, errors.New("not a ClientHello") + } + handshakeLen := int(data[1])<<16 | int(data[2])<<8 | int(data[3]) + if len(data) != 4+handshakeLen { + return 0, 0, 0, io.ErrUnexpectedEOF + } + + parsePos := 4 + // Skip protocol version (2 bytes) + if parsePos+2 > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + parsePos += 2 + // skip random (32 bytes) + if parsePos+32 > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + parsePos += 32 + // session ID + if parsePos+1 > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + sessionIDLen := int(data[parsePos]) + parsePos++ + if parsePos+sessionIDLen > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + parsePos += sessionIDLen + // cipher suites + if parsePos+2 > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + cipherSuitesLen := int(binary.BigEndian.Uint16(data[parsePos:])) + parsePos += 2 + if parsePos+cipherSuitesLen > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + parsePos += cipherSuitesLen + // compression methods + if parsePos+1 > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + compressionMethodsLen := int(data[parsePos]) + parsePos++ + if parsePos+compressionMethodsLen > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + parsePos += compressionMethodsLen + + // extensions + if parsePos+2 > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + extensionsLen := int(binary.BigEndian.Uint16(data[parsePos:])) + parsePos += 2 + if parsePos+extensionsLen > len(data) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + extensionsStart := parsePos + extensions := data[extensionsStart : extensionsStart+extensionsLen] + + // parse extensions + var extPos int + sniPos = -1 + echPos = -1 + for extPos+4 <= extensionsLen { + extType := binary.BigEndian.Uint16(extensions[extPos:]) + extLen := int(binary.BigEndian.Uint16(extensions[extPos+2:])) + if extPos+4+extLen > extensionsLen { + return 0, 0, 0, io.ErrUnexpectedEOF + } + switch extType { + case extTypeSNI: + if sniPos != -1 { + return 0, 0, 0, errors.New("multiple SNI extensions") + } + sniData := extensions[extPos+4 : extPos+4+extLen] + if len(sniData) < 2 { + return 0, 0, 0, io.ErrUnexpectedEOF + } + nameListLen := int(binary.BigEndian.Uint16(sniData)) + if len(sniData) != 2+nameListLen { + return 0, 0, 0, io.ErrUnexpectedEOF + } + listPos := 2 + for listPos+3 <= nameListLen+2 { + nameType := sniData[listPos] + sniLen = int(binary.BigEndian.Uint16(sniData[listPos+1:])) + if listPos+3+sniLen > len(sniData) { + return 0, 0, 0, io.ErrUnexpectedEOF + } + if nameType == 0 { // host_name + sniPos = extensionsStart + extPos + 4 + listPos + 3 + break // stop after first host_name + } + listPos += 3 + sniLen + } + if sniPos == 0 { + return 0, 0, 0, errors.New("SNI host_name not found") + } + case extTypeECH: + if echPos != -1 { + return 0, 0, 0, errors.New("multiple ECH extensions") + } + echPos = extensionsStart + extPos + } + extPos += 4 + extLen + if sniPos != -1 && echPos != -1 { + break + } + } + return sniPos, sniLen, echPos, nil +} diff --git a/third_party/quic-go/sni_test.go b/third_party/quic-go/sni_test.go new file mode 100644 index 0000000..b57324e --- /dev/null +++ b/third_party/quic-go/sni_test.go @@ -0,0 +1,301 @@ +package quic + +import ( + "context" + "crypto/ecdh" + "crypto/rand" + "crypto/tls" + "encoding/binary" + "fmt" + "io" + mrand "math/rand/v2" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/testdata" + "golang.org/x/crypto/cryptobyte" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func checkClientHello(clientHello []byte) error { + conn := tls.QUICServer(&tls.QUICConfig{ + TLSConfig: testdata.GetTLSConfig(), + }) + if err := conn.Start(context.Background()); err != nil { + return err + } + defer conn.Close() + return conn.HandleData(tls.QUICEncryptionLevelInitial, clientHello) +} + +func getClientHello(serverName string) ([]byte, error) { + c := tls.QUICClient(&tls.QUICConfig{ + TLSConfig: &tls.Config{ + ServerName: serverName, + MinVersion: tls.VersionTLS13, + InsecureSkipVerify: serverName == "", + // disable post-quantum curves + CurvePreferences: []tls.CurveID{tls.CurveP256}, + }, + }) + b := make([]byte, mrand.IntN(200)) + rand.Read(b) + c.SetTransportParameters(b) + if err := c.Start(context.Background()); err != nil { + return nil, err + } + + ev := c.NextEvent() + if ev.Kind != tls.QUICWriteData { + return nil, fmt.Errorf("expected QUICWriteData event, got %v", ev.Kind) + } + if err := checkClientHello(ev.Data); err != nil { + return nil, err + } + return ev.Data, nil +} + +func getClientHelloWithECH(serverName string) ([]byte, error) { + // various constants from the standard library's (internal) hpke package + const ( + DHKEM_X25519_HKDF_SHA256 = 0x20 + KDF_HKDF_SHA256 = 1 + AEAD_AES_128_GCM = 1 + ) + + marshalECHConfig := func(id uint8, pubKey []byte, publicName string, maxNameLen uint8) []byte { + builder := cryptobyte.NewBuilder(nil) + builder.AddUint16(extTypeECH) + builder.AddUint16LengthPrefixed(func(builder *cryptobyte.Builder) { + builder.AddUint8(id) + builder.AddUint16(DHKEM_X25519_HKDF_SHA256) + builder.AddUint16LengthPrefixed(func(builder *cryptobyte.Builder) { builder.AddBytes(pubKey) }) + builder.AddUint16LengthPrefixed(func(builder *cryptobyte.Builder) { + builder.AddUint16(KDF_HKDF_SHA256) + builder.AddUint16(AEAD_AES_128_GCM) + }) + builder.AddUint8(maxNameLen) + builder.AddUint8LengthPrefixed(func(builder *cryptobyte.Builder) { + builder.AddBytes([]byte(publicName)) + }) + builder.AddUint16(0) // extensions + }) + + return builder.BytesOrPanic() + } + + echKey, _ := ecdh.X25519().GenerateKey(rand.Reader) + echConfig := marshalECHConfig(42, echKey.PublicKey().Bytes(), serverName, 32) + + builder := cryptobyte.NewBuilder(nil) + builder.AddUint16LengthPrefixed(func(builder *cryptobyte.Builder) { builder.AddBytes(echConfig) }) + + c := tls.QUICClient(&tls.QUICConfig{ + TLSConfig: &tls.Config{ + ServerName: serverName, + MinVersion: tls.VersionTLS13, + EncryptedClientHelloConfigList: builder.BytesOrPanic(), + InsecureSkipVerify: serverName == "", + // disable post-quantum curves + CurvePreferences: []tls.CurveID{tls.CurveP256}, + }, + }) + b := make([]byte, mrand.IntN(200)) + rand.Read(b) + c.SetTransportParameters(b) + if err := c.Start(context.Background()); err != nil { + return nil, err + } + + ev := c.NextEvent() + if ev.Kind != tls.QUICWriteData { + return nil, fmt.Errorf("expected QUICWriteData event, got %v", ev.Kind) + } + if err := checkClientHello(ev.Data); err != nil { + return nil, err + } + return ev.Data, nil +} + +// shuffleClientHelloExtensions takes a TLS 1.3 ClientHello message (without the record layer) +// and returns a new ClientHello with its extensions shuffled. Returns nil if the input is invalid. +func shuffleClientHelloExtensions(t testing.TB, clientHello []byte) []byte { + t.Helper() + + // Basic validation: ensure minimum length and correct handshake type (0x01 for ClientHello) + if len(clientHello) < 4 || clientHello[0] != 0x01 { + t.Fatalf("not a ClientHello") + } + + // Extract the 3-byte length (24-bit integer) and validate total length + length := uint32(clientHello[1])<<16 | uint32(clientHello[2])<<8 | uint32(clientHello[3]) + require.Equal(t, 4+int(length), len(clientHello)) + + // Body is everything after type and length + body := clientHello[4 : 4+length] + var pos int + // Parse fixed and variable-length fields to reach extensions + require.Greater(t, len(body), pos+2) // protocol version: 2 bytes + pos += 2 + require.Greater(t, len(body), pos+32) // random: 32 bytes + pos += 32 + require.Greater(t, len(body), pos+1) // session ID length: 1 byte + sessionIDLen := int(body[pos]) + pos += 1 + require.Greater(t, len(body), pos+sessionIDLen) // session ID data + pos += sessionIDLen + require.Greater(t, len(body), pos+2) // cipher suites length: 2 bytes + cipherSuitesLen := int(body[pos])<<8 | int(body[pos+1]) + pos += 2 + require.Greater(t, len(body), pos+cipherSuitesLen) // cipher suites data + pos += cipherSuitesLen + require.Greater(t, len(body), pos+1) // compression methods length: 1 byte + compressionMethodsLen := int(body[pos]) + pos += 1 + require.Greater(t, len(body), pos+compressionMethodsLen) // compression methods data + pos += compressionMethodsLen + + // Extensions: 2 bytes total length + data (may be absent) + if pos+2 > len(body) { + // No extensions present; return original + return clientHello + } + extensionsLen := int(body[pos])<<8 | int(body[pos+1]) + pos += 2 + require.Equal(t, pos+extensionsLen, len(body)) // extensions length doesn't match remaining data + extensionsData := body[pos : pos+extensionsLen] + + // parse extensions into a slice of byte slices + var extensions [][]byte + var extPos int + for extPos < extensionsLen { + require.Greater(t, extensionsLen, extPos+4) // type and length + extLen := int(extensionsData[extPos+2])<<8 | int(extensionsData[extPos+3]) + require.LessOrEqual(t, extPos+4+extLen, extensionsLen) // extension exceeds total length + // extract entire extension (type: 2 bytes, length: 2 bytes, data) + extData := extensionsData[extPos : extPos+4+extLen] + extensions = append(extensions, extData) + extPos += 4 + extLen + } + + // shuffle extensions using a proper random source + mrand.Shuffle(len(extensions), func(i, j int) { + extensions[i], extensions[j] = extensions[j], extensions[i] + }) + + // reconstruct extensions data + var newExtensionsData []byte + for _, ext := range extensions { + newExtensionsData = append(newExtensionsData, ext...) + } + + // reconstruct body: prefix (up to and including extensions length) + shuffled extensions + prefix := body[:pos] + newBody := append(prefix, newExtensionsData...) + // reconstruct ClientHello: type (0x01) + original length + new body + newClientHello := []byte{0x01} + lengthBytes := clientHello[1:4] // length unchanged since only extensions are shuffled + newClientHello = append(newClientHello, lengthBytes...) + newClientHello = append(newClientHello, newBody...) + // check that it's actually valid + if err := checkClientHello(newClientHello); err != nil { + t.Fatalf("invalid ClientHello: %v", err) + } + return newClientHello +} + +func TestFindSNI(t *testing.T) { + t.Run("without SNI", func(t *testing.T) { + testFindSNI(t, "") + }) + t.Run("without subdomain", func(t *testing.T) { + testFindSNI(t, "quic-go.net") + }) + t.Run("with subdomain", func(t *testing.T) { + testFindSNI(t, "sub.do.ma.in.quic-go.net") + }) +} + +func testFindSNI(t *testing.T, serverName string) { + clientHello, err := getClientHello(serverName) + require.NoError(t, err) + sniPos, sniLen, echPos, err := findSNIAndECH(clientHello) + require.NoError(t, err) + assert.Equal(t, -1, echPos) + if serverName == "" { + require.Equal(t, -1, sniPos) + return + } + assert.Equal(t, len(serverName), sniLen) + require.NotEqual(t, -1, sniPos) + require.Equal(t, serverName, string(clientHello[sniPos:sniPos+sniLen])) + + // incomplete ClientHellos result in an io.ErrUnexpectedEOF + for i := range clientHello { + _, _, _, err := findSNIAndECH(clientHello[:i]) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + } +} + +func TestFindSNIWithECH(t *testing.T) { + const serverName = "public.example" + clientHello, err := getClientHelloWithECH(serverName) + require.NoError(t, err) + clientHello = shuffleClientHelloExtensions(t, clientHello) + sniPos, sniLen, echPos, err := findSNIAndECH(clientHello) + require.NoError(t, err) + require.NotEqual(t, -1, echPos) + require.Equal(t, uint16(extTypeECH), binary.BigEndian.Uint16(clientHello[echPos:echPos+2])) + assert.Equal(t, len(serverName), sniLen) + require.NotEqual(t, -1, sniPos) + require.Equal(t, serverName, string(clientHello[sniPos:sniPos+sniLen])) + + for i := range clientHello { + _, _, _, err := findSNIAndECH(clientHello[:i]) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + } +} + +// findSNI is never run with attacker-controlled inputs (other than the session ticket), +// so this is not a high-value target to begin with, +// and doesn't need to be run in ClusterFuzz. +// It's still useful to find potential corner cases in the parser. +func FuzzFindSNI(f *testing.F) { + addSeed := func(serverName string, maxSize int) { + ch, err := getClientHello(serverName) + require.NoError(f, err) + f.Add(ch, maxSize) + } + addSeedWithECH := func(serverName string, maxSize int) { + ch, err := getClientHelloWithECH(serverName) + require.NoError(f, err) + f.Add(ch, maxSize) + } + + addSeed("", 10) + addSeed("google.com", 20) + addSeed("sub.do.ma.in.quic-go.net", 30) + addSeedWithECH("quic-go.net", 40) + + f.Fuzz(func(t *testing.T, data []byte, maxSize int) { + cs := newInitialCryptoStream(true, false) + if _, err := cs.Write(data); err != nil { + return + } + segments := make(map[protocol.ByteCount][]byte) + if !cs.HasData() { // incomplete ClientHello + return + } + for cs.HasData() { + f := cs.PopCryptoFrame(5 + protocol.ByteCount(maxSize)) + if f == nil { + return + } + segments[f.Offset] = f.Data + } + reassembled := reassembleCryptoData(t, segments) + require.Equal(t, data, reassembled) + }) +} diff --git a/third_party/quic-go/stateless_reset.go b/third_party/quic-go/stateless_reset.go new file mode 100644 index 0000000..f1d4fbd --- /dev/null +++ b/third_party/quic-go/stateless_reset.go @@ -0,0 +1,42 @@ +package quic + +import ( + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "hash" + "sync" + + "github.com/apernet/quic-go/internal/protocol" +) + +type statelessResetter struct { + mx sync.Mutex + h hash.Hash +} + +// newStatelessRetter creates a new stateless reset generator. +// It is valid to use a nil key. In that case, a random key will be used. +// This makes is impossible for on-path attackers to shut down established connections. +func newStatelessResetter(key *StatelessResetKey) *statelessResetter { + var h hash.Hash + if key != nil { + h = hmac.New(sha256.New, key[:]) + } else { + b := make([]byte, 32) + _, _ = rand.Read(b) + h = hmac.New(sha256.New, b) + } + return &statelessResetter{h: h} +} + +func (r *statelessResetter) GetStatelessResetToken(connID protocol.ConnectionID) protocol.StatelessResetToken { + r.mx.Lock() + defer r.mx.Unlock() + + var token protocol.StatelessResetToken + r.h.Write(connID.Bytes()) + copy(token[:], r.h.Sum(nil)) + r.h.Reset() + return token +} diff --git a/third_party/quic-go/stateless_reset_test.go b/third_party/quic-go/stateless_reset_test.go new file mode 100644 index 0000000..ab5653b --- /dev/null +++ b/third_party/quic-go/stateless_reset_test.go @@ -0,0 +1,42 @@ +package quic + +import ( + "crypto/rand" + "testing" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/stretchr/testify/require" +) + +func TestStatelessResetter(t *testing.T) { + t.Run("no key", func(t *testing.T) { + r1 := newStatelessResetter(nil) + r2 := newStatelessResetter(nil) + for range 100 { + b := make([]byte, 15) + rand.Read(b) + connID := protocol.ParseConnectionID(b) + t1 := r1.GetStatelessResetToken(connID) + t2 := r2.GetStatelessResetToken(connID) + require.NotZero(t, t1) + require.NotZero(t, t2) + require.NotEqual(t, t1, t2) + } + }) + + t.Run("with key", func(t *testing.T) { + var key StatelessResetKey + rand.Read(key[:]) + m := newStatelessResetter(&key) + b := make([]byte, 8) + rand.Read(b) + connID := protocol.ParseConnectionID(b) + token := m.GetStatelessResetToken(connID) + require.NotZero(t, token) + require.Equal(t, token, m.GetStatelessResetToken(connID)) + // generate a new connection ID + rand.Read(b) + connID2 := protocol.ParseConnectionID(b) + require.NotEqual(t, token, m.GetStatelessResetToken(connID2)) + }) +} diff --git a/third_party/quic-go/stream.go b/third_party/quic-go/stream.go new file mode 100644 index 0000000..a3155cf --- /dev/null +++ b/third_party/quic-go/stream.go @@ -0,0 +1,253 @@ +package quic + +import ( + "context" + "net" + "os" + "sync" + "time" + + "github.com/apernet/quic-go/internal/ackhandler" + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +type deadlineError struct{} + +func (deadlineError) Error() string { return "deadline exceeded" } +func (deadlineError) Temporary() bool { return true } +func (deadlineError) Timeout() bool { return true } +func (deadlineError) Unwrap() error { return os.ErrDeadlineExceeded } + +var errDeadline net.Error = &deadlineError{} + +// The streamSender is notified by the stream about various events. +type streamSender interface { + onHasConnectionData() + onHasStreamData(protocol.StreamID, *SendStream) + onHasStreamControlFrame(protocol.StreamID, streamControlFrameGetter) + // must be called without holding the mutex that is acquired by closeForShutdown + onStreamCompleted(protocol.StreamID) +} + +// Each of the both stream halves gets its own uniStreamSender. +// This is necessary in order to keep track when both halves have been completed. +type uniStreamSender struct { + streamSender + onStreamCompletedImpl func() + onHasStreamControlFrameImpl func(protocol.StreamID, streamControlFrameGetter) +} + +func (s *uniStreamSender) onHasStreamData(id protocol.StreamID, str *SendStream) { + s.streamSender.onHasStreamData(id, str) +} +func (s *uniStreamSender) onStreamCompleted(protocol.StreamID) { s.onStreamCompletedImpl() } +func (s *uniStreamSender) onHasStreamControlFrame(id protocol.StreamID, str streamControlFrameGetter) { + s.onHasStreamControlFrameImpl(id, str) +} + +var _ streamSender = &uniStreamSender{} + +type Stream struct { + receiveStr *ReceiveStream + sendStr *SendStream + + completedMutex sync.Mutex + sender streamSender + receiveStreamCompleted bool + sendStreamCompleted bool +} + +var ( + _ outgoingStream = &Stream{} + _ sendStreamFrameHandler = &Stream{} + _ receiveStreamFrameHandler = &Stream{} +) + +// newStream creates a new Stream +func newStream( + ctx context.Context, + streamID protocol.StreamID, + sender streamSender, + flowController *streamFlowController, + supportsResetStreamAt bool, +) *Stream { + s := &Stream{sender: sender} + senderForSendStream := &uniStreamSender{ + streamSender: sender, + onStreamCompletedImpl: func() { + s.completedMutex.Lock() + s.sendStreamCompleted = true + s.checkIfCompleted() + s.completedMutex.Unlock() + }, + onHasStreamControlFrameImpl: func(id protocol.StreamID, str streamControlFrameGetter) { + sender.onHasStreamControlFrame(streamID, s) + }, + } + s.sendStr = newSendStream(ctx, streamID, senderForSendStream, flowController, supportsResetStreamAt) + senderForReceiveStream := &uniStreamSender{ + streamSender: sender, + onStreamCompletedImpl: func() { + s.completedMutex.Lock() + s.receiveStreamCompleted = true + s.checkIfCompleted() + s.completedMutex.Unlock() + }, + onHasStreamControlFrameImpl: func(id protocol.StreamID, str streamControlFrameGetter) { + sender.onHasStreamControlFrame(streamID, s) + }, + } + s.receiveStr = newReceiveStream(streamID, senderForReceiveStream, flowController) + return s +} + +// StreamID returns the stream ID. +func (s *Stream) StreamID() StreamID { + // the result is same for receiveStream and sendStream + return s.sendStr.StreamID() +} + +// Read reads data from the stream. +// Read can be made to time out using [Stream.SetReadDeadline] and [Stream.SetDeadline]. +// If the stream was canceled, the error is a [StreamError]. +func (s *Stream) Read(p []byte) (int, error) { + return s.receiveStr.Read(p) +} + +// Peek fills b with stream data, without consuming the stream data. +// It blocks until len(b) bytes are available, or an error occurs. +// It respects the stream deadline set by SetReadDeadline. +// If the stream ends before len(b) bytes are available, +// it returns the number of bytes peeked along with io.EOF. +func (s *Stream) Peek(b []byte) (int, error) { + return s.receiveStr.Peek(b) +} + +// Write writes data to the stream. +// Write can be made to time out using [Stream.SetWriteDeadline] or [Stream.SetDeadline]. +// If the stream was canceled, the error is a [StreamError]. +func (s *Stream) Write(p []byte) (int, error) { + return s.sendStr.Write(p) +} + +// WriteWithLimit writes data to the stream, subject to an additional send limit. +// See [SendStream.WriteWithLimit] for more details. +func (s *Stream) WriteWithLimit(p []byte, limiter func(maxBytes int) int) (int, error) { + return s.sendStr.WriteWithLimit(p, limiter) +} + +// TryWriteAll writes data to the stream if it can be queued immediately. +// See [SendStream.TryWriteAll] for more details. +func (s *Stream) TryWriteAll(p []byte) error { + return s.sendStr.TryWriteAll(p) +} + +// SetReliableBoundary marks the data written to this stream so far as reliable. +// It is valid to call this function multiple times, thereby increasing the reliable size. +// It only has an effect if the peer enabled support for the RESET_STREAM_AT extension, +// otherwise, it is a no-op. +func (s *Stream) SetReliableBoundary() { + s.sendStr.SetReliableBoundary() +} + +// CancelWrite aborts sending on this stream. +// See [SendStream.CancelWrite] for more details. +func (s *Stream) CancelWrite(errorCode StreamErrorCode) { + s.sendStr.CancelWrite(errorCode) +} + +// CancelRead aborts receiving on this stream. +// See [ReceiveStream.CancelRead] for more details. +func (s *Stream) CancelRead(errorCode StreamErrorCode) { + s.receiveStr.CancelRead(errorCode) +} + +// SetReceiveFinalSizeCallback sets a callback that is called when the receive side's final size is known. +// See [ReceiveStream.SetReceiveFinalSizeCallback] for more details. +// Most applications don't need this. It is mainly useful for protocol layers +// that need exact stream final sizes, such as WebTransport flow control accounting. +func (s *Stream) SetReceiveFinalSizeCallback(callback func(int64)) { + s.receiveStr.SetReceiveFinalSizeCallback(callback) +} + +// The Context is canceled as soon as the write-side of the stream is closed. +// See [SendStream.Context] for more details. +func (s *Stream) Context() context.Context { + return s.sendStr.Context() +} + +// Close closes the send-direction of the stream. +// It does not close the receive-direction of the stream. +func (s *Stream) Close() error { + return s.sendStr.Close() +} + +func (s *Stream) handleResetStreamFrame(frame *wire.ResetStreamFrame, rcvTime monotime.Time) error { + return s.receiveStr.handleResetStreamFrame(frame, rcvTime) +} + +func (s *Stream) handleStreamFrame(frame *wire.StreamFrame, rcvTime monotime.Time) error { + return s.receiveStr.handleStreamFrame(frame, rcvTime) +} + +func (s *Stream) handleStopSendingFrame(frame *wire.StopSendingFrame) { + s.sendStr.handleStopSendingFrame(frame) +} + +func (s *Stream) updateSendWindow(limit protocol.ByteCount) { + s.sendStr.updateSendWindow(limit) +} + +func (s *Stream) enableResetStreamAt() { + s.sendStr.enableResetStreamAt() +} + +func (s *Stream) popStreamFrame(maxBytes protocol.ByteCount, v protocol.Version) (_ ackhandler.StreamFrame, _ *wire.StreamDataBlockedFrame, hasMore bool) { + return s.sendStr.popStreamFrame(maxBytes, v) +} + +func (s *Stream) getControlFrame(now monotime.Time) (_ ackhandler.Frame, ok, hasMore bool) { + f, ok, _ := s.sendStr.getControlFrame(now) + if ok { + return f, true, true + } + return s.receiveStr.getControlFrame(now) +} + +// SetReadDeadline sets the deadline for future Read calls. +// See [ReceiveStream.SetReadDeadline] for more details. +func (s *Stream) SetReadDeadline(t time.Time) error { + return s.receiveStr.SetReadDeadline(t) +} + +// SetWriteDeadline sets the deadline for future Write calls. +// See [SendStream.SetWriteDeadline] for more details. +func (s *Stream) SetWriteDeadline(t time.Time) error { + return s.sendStr.SetWriteDeadline(t) +} + +// SetDeadline sets the read and write deadlines associated with the stream. +// It is equivalent to calling both SetReadDeadline and SetWriteDeadline. +func (s *Stream) SetDeadline(t time.Time) error { + _ = s.receiveStr.SetReadDeadline(t) // SetReadDeadline never errors + _ = s.sendStr.SetWriteDeadline(t) // SetWriteDeadline never errors + return nil +} + +// CloseForShutdown closes a stream abruptly. +// It makes Read and Write unblock (and return the error) immediately. +// The peer will NOT be informed about this: the stream is closed without sending a FIN or RST. +func (s *Stream) closeForShutdown(err error) { + s.sendStr.closeForShutdown(err) + s.receiveStr.closeForShutdown(err) +} + +// checkIfCompleted is called from the uniStreamSender, when one of the stream halves is completed. +// It makes sure that the onStreamCompleted callback is only called if both receive and send side have completed. +func (s *Stream) checkIfCompleted() { + if s.sendStreamCompleted && s.receiveStreamCompleted { + s.sender.onStreamCompleted(s.StreamID()) + } +} diff --git a/third_party/quic-go/stream_test.go b/third_party/quic-go/stream_test.go new file mode 100644 index 0000000..06bd69b --- /dev/null +++ b/third_party/quic-go/stream_test.go @@ -0,0 +1,109 @@ +package quic + +import ( + "context" + "io" + "os" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func TestStreamDeadlines(t *testing.T) { + const streamID protocol.StreamID = 1337 + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + fc := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + str := newStream(context.Background(), streamID, mockSender, fc, false) + + // SetDeadline sets both read and write deadlines + str.SetDeadline(time.Now().Add(-time.Second)) + n, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Zero(t, n) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar")}, monotime.Now())) + n, err = (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(make([]byte, 6)) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.Zero(t, n) +} + +func TestStreamReceiveFinalSizeCallback(t *testing.T) { + const streamID protocol.StreamID = 1337 + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + fc := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + str := newStream(context.Background(), streamID, mockSender, fc, false) + + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{Data: []byte("foobar"), Fin: true}, monotime.Now())) + + var size int64 + str.SetReceiveFinalSizeCallback(func(s int64) { size = s }) + require.EqualValues(t, 6, size) +} + +func TestStreamCompletion(t *testing.T) { + completeReadSide := func( + t *testing.T, + str *Stream, + mockCtrl *gomock.Controller, + ) { + t.Helper() + require.NoError(t, str.handleStreamFrame(&wire.StreamFrame{ + StreamID: str.StreamID(), + Data: []byte("foobar"), + Fin: true, + }, monotime.Now())) + _, err := (&readerWithTimeout{Reader: str, Timeout: time.Second}).Read(make([]byte, 6)) + require.ErrorIs(t, err, io.EOF) + require.True(t, mockCtrl.Satisfied()) + } + + completeWriteSide := func( + t *testing.T, + str *Stream, + mockCtrl *gomock.Controller, + mockSender *MockStreamSender, + ) { + t.Helper() + mockSender.EXPECT().onHasStreamData(str.StreamID(), gomock.Any()).Times(2) + _, err := (&writerWithTimeout{Writer: str, Timeout: time.Second}).Write([]byte("foobar")) + require.NoError(t, err) + require.NoError(t, str.Close()) + f, _, _ := str.popStreamFrame(protocol.MaxByteCount, protocol.Version1) + require.NotNil(t, f.Frame) + require.True(t, f.Frame.Fin) + f.Handler.OnAcked(f.Frame) + require.True(t, mockCtrl.Satisfied()) + } + + const streamID protocol.StreamID = 1337 + + t.Run("first read, then write", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + fc := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + str := newStream(context.Background(), streamID, mockSender, fc, false) + + completeReadSide(t, str, mockCtrl) + mockSender.EXPECT().onStreamCompleted(streamID) + completeWriteSide(t, str, mockCtrl, mockSender) + }) + + t.Run("first write, then read", func(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + fc := newTestStreamFlowControllerWithSendWindow(streamID, protocol.MaxByteCount) + str := newStream(context.Background(), streamID, mockSender, fc, false) + + completeWriteSide(t, str, mockCtrl, mockSender) + mockSender.EXPECT().onStreamCompleted(streamID) + completeReadSide(t, str, mockCtrl) + }) +} diff --git a/third_party/quic-go/streams_map.go b/third_party/quic-go/streams_map.go new file mode 100644 index 0000000..f9b3e1a --- /dev/null +++ b/third_party/quic-go/streams_map.go @@ -0,0 +1,356 @@ +package quic + +import ( + "context" + "fmt" + "sync" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +// StreamLimitReachedError is returned from Conn.OpenStream and Conn.OpenUniStream +// when it is not possible to open a new stream because the number of opens streams reached +// the peer's stream limit. +type StreamLimitReachedError struct{} + +func (e StreamLimitReachedError) Error() string { return "too many open streams" } + +type streamsMap struct { + ctx context.Context // not used for cancellations, but carries the values associated with the connection + perspective protocol.Perspective + + maxIncomingBidiStreams uint64 + maxIncomingUniStreams uint64 + + sender streamSender + queueControlFrame func(wire.Frame) + newFlowController func(protocol.StreamID) *streamFlowController + + mutex sync.Mutex + outgoingBidiStreams *outgoingStreamsMap[*Stream] + outgoingUniStreams *outgoingStreamsMap[*SendStream] + incomingBidiStreams *incomingStreamsMap[*Stream] + incomingUniStreams *incomingStreamsMap[*ReceiveStream] + reset bool + supportsResetStreamAt bool +} + +func newStreamsMap( + ctx context.Context, + sender streamSender, + queueControlFrame func(wire.Frame), + newFlowController func(protocol.StreamID) *streamFlowController, + maxIncomingBidiStreams uint64, + maxIncomingUniStreams uint64, + perspective protocol.Perspective, +) *streamsMap { + m := &streamsMap{ + ctx: ctx, + perspective: perspective, + queueControlFrame: queueControlFrame, + newFlowController: newFlowController, + maxIncomingBidiStreams: maxIncomingBidiStreams, + maxIncomingUniStreams: maxIncomingUniStreams, + sender: sender, + } + m.initMaps() + return m +} + +func (m *streamsMap) initMaps() { + m.outgoingBidiStreams = newOutgoingStreamsMap( + protocol.StreamTypeBidi, + func(id protocol.StreamID) *Stream { + return newStream(m.ctx, id, m.sender, m.newFlowController(id), m.supportsResetStreamAt) + }, + m.queueControlFrame, + m.perspective, + ) + m.incomingBidiStreams = newIncomingStreamsMap( + protocol.StreamTypeBidi, + func(id protocol.StreamID) *Stream { + return newStream(m.ctx, id, m.sender, m.newFlowController(id), m.supportsResetStreamAt) + }, + m.maxIncomingBidiStreams, + m.queueControlFrame, + m.perspective, + ) + m.outgoingUniStreams = newOutgoingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *SendStream { + return newSendStream(m.ctx, id, m.sender, m.newFlowController(id), m.supportsResetStreamAt) + }, + m.queueControlFrame, + m.perspective, + ) + m.incomingUniStreams = newIncomingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *ReceiveStream { + return newReceiveStream(id, m.sender, m.newFlowController(id)) + }, + m.maxIncomingUniStreams, + m.queueControlFrame, + m.perspective, + ) +} + +func (m *streamsMap) OpenStream() (*Stream, error) { + m.mutex.Lock() + reset := m.reset + mm := m.outgoingBidiStreams + m.mutex.Unlock() + if reset { + return nil, Err0RTTRejected + } + return mm.OpenStream() +} + +func (m *streamsMap) OpenStreamSync(ctx context.Context) (*Stream, error) { + m.mutex.Lock() + reset := m.reset + mm := m.outgoingBidiStreams + m.mutex.Unlock() + if reset { + return nil, Err0RTTRejected + } + return mm.OpenStreamSync(ctx) +} + +func (m *streamsMap) OpenUniStream() (*SendStream, error) { + m.mutex.Lock() + reset := m.reset + mm := m.outgoingUniStreams + m.mutex.Unlock() + if reset { + return nil, Err0RTTRejected + } + return mm.OpenStream() +} + +func (m *streamsMap) OpenUniStreamSync(ctx context.Context) (*SendStream, error) { + m.mutex.Lock() + reset := m.reset + mm := m.outgoingUniStreams + m.mutex.Unlock() + if reset { + return nil, Err0RTTRejected + } + return mm.OpenStreamSync(ctx) +} + +func (m *streamsMap) AcceptStream(ctx context.Context) (*Stream, error) { + m.mutex.Lock() + reset := m.reset + mm := m.incomingBidiStreams + m.mutex.Unlock() + if reset { + return nil, Err0RTTRejected + } + return mm.AcceptStream(ctx) +} + +func (m *streamsMap) AcceptUniStream(ctx context.Context) (*ReceiveStream, error) { + m.mutex.Lock() + reset := m.reset + mm := m.incomingUniStreams + m.mutex.Unlock() + if reset { + return nil, Err0RTTRejected + } + return mm.AcceptStream(ctx) +} + +func (m *streamsMap) DeleteStream(id protocol.StreamID) error { + switch protocol.StreamTypeOf(id) { + case protocol.StreamTypeUni: + if protocol.StreamInitiator(id) == m.perspective { + return m.outgoingUniStreams.DeleteStream(id) + } + return m.incomingUniStreams.DeleteStream(id) + case protocol.StreamTypeBidi: + if protocol.StreamInitiator(id) == m.perspective { + return m.outgoingBidiStreams.DeleteStream(id) + } + return m.incomingBidiStreams.DeleteStream(id) + } + panic("") +} + +func (m *streamsMap) HandleMaxStreamsFrame(f *wire.MaxStreamsFrame) { + switch f.Type { + case protocol.StreamTypeUni: + m.outgoingUniStreams.SetMaxStream(f.MaxStreamNum.StreamID(protocol.StreamTypeUni, m.perspective)) + case protocol.StreamTypeBidi: + m.outgoingBidiStreams.SetMaxStream(f.MaxStreamNum.StreamID(protocol.StreamTypeBidi, m.perspective)) + } +} + +type sendStreamFrameHandler interface { + updateSendWindow(protocol.ByteCount) + handleStopSendingFrame(*wire.StopSendingFrame) +} + +func (m *streamsMap) getSendStream(id protocol.StreamID) (sendStreamFrameHandler, error) { + switch protocol.StreamTypeOf(id) { + case protocol.StreamTypeUni: + if protocol.StreamInitiator(id) != m.perspective { + // an outgoing unidirectional stream is a send stream, not a receive stream + return nil, &qerr.TransportError{ + ErrorCode: qerr.StreamStateError, + ErrorMessage: fmt.Sprintf("invalid frame for send stream %d", id), + } + } + str, err := m.outgoingUniStreams.GetStream(id) + if str == nil || err != nil { + return nil, err + } + return str, nil + case protocol.StreamTypeBidi: + if protocol.StreamInitiator(id) == m.perspective { + str, err := m.outgoingBidiStreams.GetStream(id) + if str == nil || err != nil { + return nil, err + } + return str, nil + } + str, err := m.incomingBidiStreams.GetOrOpenStream(id) + if str == nil || err != nil { + return nil, err + } + return str, nil + } + panic("unreachable") +} + +func (m *streamsMap) HandleMaxStreamDataFrame(f *wire.MaxStreamDataFrame) error { + str, err := m.getSendStream(f.StreamID) + if err != nil { + return err + } + if str == nil { // stream already deleted + return nil + } + str.updateSendWindow(f.MaximumStreamData) + return nil +} + +func (m *streamsMap) HandleStopSendingFrame(f *wire.StopSendingFrame) error { + str, err := m.getSendStream(f.StreamID) + if err != nil { + return err + } + if str == nil { // stream already deleted + return nil + } + str.handleStopSendingFrame(f) + return nil +} + +type receiveStreamFrameHandler interface { + handleResetStreamFrame(*wire.ResetStreamFrame, monotime.Time) error + handleStreamFrame(*wire.StreamFrame, monotime.Time) error +} + +func (m *streamsMap) getReceiveStream(id protocol.StreamID) (receiveStreamFrameHandler, error) { + switch protocol.StreamTypeOf(id) { + case protocol.StreamTypeUni: + // an outgoing unidirectional stream is a send stream, not a receive stream + if protocol.StreamInitiator(id) == m.perspective { + return nil, &qerr.TransportError{ + ErrorCode: qerr.StreamStateError, + ErrorMessage: fmt.Sprintf("invalid frame for receive stream %d", id), + } + } + str, err := m.incomingUniStreams.GetOrOpenStream(id) + if err != nil || str == nil { + return nil, err + } + return str, nil + case protocol.StreamTypeBidi: + var str *Stream + var err error + if protocol.StreamInitiator(id) == m.perspective { + str, err = m.outgoingBidiStreams.GetStream(id) + } else { + str, err = m.incomingBidiStreams.GetOrOpenStream(id) + } + if str == nil || err != nil { + return nil, err + } + return str, nil + } + panic("unreachable") +} + +func (m *streamsMap) HandleStreamDataBlockedFrame(f *wire.StreamDataBlockedFrame) error { + if _, err := m.getReceiveStream(f.StreamID); err != nil { + return err + } + // We don't need to do anything in response to a STREAM_DATA_BLOCKED frame, + // but we need to make sure that the stream ID is valid. + return nil // we don't need to do anything in response to a STREAM_DATA_BLOCKED frame +} + +func (m *streamsMap) HandleResetStreamFrame(f *wire.ResetStreamFrame, rcvTime monotime.Time) error { + str, err := m.getReceiveStream(f.StreamID) + if err != nil { + return err + } + if str == nil { // stream already deleted + return nil + } + return str.handleResetStreamFrame(f, rcvTime) +} + +func (m *streamsMap) HandleStreamFrame(f *wire.StreamFrame, rcvTime monotime.Time) error { + str, err := m.getReceiveStream(f.StreamID) + if err != nil { + return err + } + if str == nil { // stream already deleted + return nil + } + return str.handleStreamFrame(f, rcvTime) +} + +func (m *streamsMap) HandleTransportParameters(p *wire.TransportParameters) { + m.supportsResetStreamAt = p.EnableResetStreamAt + if p.EnableResetStreamAt { + m.outgoingBidiStreams.EnableResetStreamAt() + m.outgoingUniStreams.EnableResetStreamAt() + } + m.outgoingBidiStreams.UpdateSendWindow(p.InitialMaxStreamDataBidiRemote) + m.outgoingBidiStreams.SetMaxStream(p.MaxBidiStreamNum.StreamID(protocol.StreamTypeBidi, m.perspective)) + m.outgoingUniStreams.UpdateSendWindow(p.InitialMaxStreamDataUni) + m.outgoingUniStreams.SetMaxStream(p.MaxUniStreamNum.StreamID(protocol.StreamTypeUni, m.perspective)) +} + +func (m *streamsMap) CloseWithError(err error) { + m.outgoingBidiStreams.CloseWithError(err) + m.outgoingUniStreams.CloseWithError(err) + m.incomingBidiStreams.CloseWithError(err) + m.incomingUniStreams.CloseWithError(err) +} + +// ResetFor0RTT resets is used when 0-RTT is rejected. In that case, the streams maps are +// 1. closed with an Err0RTTRejected, making calls to Open{Uni}Stream{Sync} / Accept{Uni}Stream return that error. +// 2. reset to their initial state, such that we can immediately process new incoming stream data. +// Afterwards, calls to Open{Uni}Stream{Sync} / Accept{Uni}Stream will continue to return the error, +// until UseResetMaps() has been called. +func (m *streamsMap) ResetFor0RTT() { + m.mutex.Lock() + defer m.mutex.Unlock() + m.reset = true + m.CloseWithError(Err0RTTRejected) + m.supportsResetStreamAt = false + m.initMaps() +} + +func (m *streamsMap) UseResetMaps() { + m.mutex.Lock() + m.reset = false + m.mutex.Unlock() +} diff --git a/third_party/quic-go/streams_map_incoming.go b/third_party/quic-go/streams_map_incoming.go new file mode 100644 index 0000000..44f128e --- /dev/null +++ b/third_party/quic-go/streams_map_incoming.go @@ -0,0 +1,209 @@ +package quic + +import ( + "context" + "fmt" + "sync" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +type incomingStream interface { + closeForShutdown(error) +} + +// When a stream is deleted before it was accepted, we can't delete it from the map immediately. +// We need to wait until the application accepts it, and delete it then. +type incomingStreamEntry[T incomingStream] struct { + stream T + shouldDelete bool +} + +type incomingStreamsMap[T incomingStream] struct { + mutex sync.RWMutex + newStreamChan chan struct{} + + streamType protocol.StreamType + streams map[protocol.StreamID]incomingStreamEntry[T] + + nextStreamToAccept protocol.StreamID // the next stream that will be returned by AcceptStream() + nextStreamToOpen protocol.StreamID // the highest stream that the peer opened + maxStream protocol.StreamID // the highest stream that the peer is allowed to open + maxNumStreams uint64 // maximum number of streams + + newStream func(protocol.StreamID) T + queueMaxStreamID func(*wire.MaxStreamsFrame) + + closeErr error +} + +func newIncomingStreamsMap[T incomingStream]( + streamType protocol.StreamType, + newStream func(protocol.StreamID) T, + maxStreams uint64, + queueControlFrame func(wire.Frame), + pers protocol.Perspective, +) *incomingStreamsMap[T] { + var nextStreamToAccept protocol.StreamID + switch { + case streamType == protocol.StreamTypeBidi && pers == protocol.PerspectiveServer: + nextStreamToAccept = protocol.FirstIncomingBidiStreamServer + case streamType == protocol.StreamTypeBidi && pers == protocol.PerspectiveClient: + nextStreamToAccept = protocol.FirstIncomingBidiStreamClient + case streamType == protocol.StreamTypeUni && pers == protocol.PerspectiveServer: + nextStreamToAccept = protocol.FirstIncomingUniStreamServer + case streamType == protocol.StreamTypeUni && pers == protocol.PerspectiveClient: + nextStreamToAccept = protocol.FirstIncomingUniStreamClient + } + return &incomingStreamsMap[T]{ + newStreamChan: make(chan struct{}, 1), + streamType: streamType, + streams: make(map[protocol.StreamID]incomingStreamEntry[T]), + maxStream: protocol.StreamNum(maxStreams).StreamID(streamType, pers.Opposite()), + maxNumStreams: maxStreams, + newStream: newStream, + nextStreamToOpen: nextStreamToAccept, + nextStreamToAccept: nextStreamToAccept, + queueMaxStreamID: func(f *wire.MaxStreamsFrame) { queueControlFrame(f) }, + } +} + +func (m *incomingStreamsMap[T]) AcceptStream(ctx context.Context) (T, error) { + // drain the newStreamChan, so we don't check the map twice if the stream doesn't exist + select { + case <-m.newStreamChan: + default: + } + + m.mutex.Lock() + + var id protocol.StreamID + var entry incomingStreamEntry[T] + for { + id = m.nextStreamToAccept + if m.closeErr != nil { + m.mutex.Unlock() + return *new(T), m.closeErr + } + var ok bool + entry, ok = m.streams[id] + if ok { + break + } + m.mutex.Unlock() + select { + case <-ctx.Done(): + return *new(T), ctx.Err() + case <-m.newStreamChan: + } + m.mutex.Lock() + } + m.nextStreamToAccept += 4 + // If this stream was completed before being accepted, we can delete it now. + if entry.shouldDelete { + if err := m.deleteStream(id); err != nil { + m.mutex.Unlock() + return *new(T), err + } + } + m.mutex.Unlock() + return entry.stream, nil +} + +func (m *incomingStreamsMap[T]) GetOrOpenStream(id protocol.StreamID) (T, error) { + m.mutex.RLock() + if id > m.maxStream { + m.mutex.RUnlock() + return *new(T), &qerr.TransportError{ + ErrorCode: qerr.StreamLimitError, + ErrorMessage: fmt.Sprintf("peer tried to open stream %d (current limit: %d)", id, m.maxStream), + } + } + // if the num is smaller than the highest we accepted + // * this stream exists in the map, and we can return it, or + // * this stream was already closed, then we can return the nil + if id < m.nextStreamToOpen { + var s T + // If the stream was already queued for deletion, and is just waiting to be accepted, don't return it. + if entry, ok := m.streams[id]; ok && !entry.shouldDelete { + s = entry.stream + } + m.mutex.RUnlock() + return s, nil + } + m.mutex.RUnlock() + + m.mutex.Lock() + // no need to check the two error conditions from above again + // * maxStream can only increase, so if the id was valid before, it definitely is valid now + // * highestStream is only modified by this function + for newNum := m.nextStreamToOpen; newNum <= id; newNum += 4 { + m.streams[newNum] = incomingStreamEntry[T]{stream: m.newStream(newNum)} + select { + case m.newStreamChan <- struct{}{}: + default: + } + } + m.nextStreamToOpen = id + 4 + entry := m.streams[id] + m.mutex.Unlock() + return entry.stream, nil +} + +func (m *incomingStreamsMap[T]) DeleteStream(id protocol.StreamID) error { + m.mutex.Lock() + defer m.mutex.Unlock() + + if err := m.deleteStream(id); err != nil { + return &qerr.TransportError{ + ErrorCode: qerr.StreamStateError, + ErrorMessage: err.Error(), + } + } + return nil +} + +func (m *incomingStreamsMap[T]) deleteStream(id protocol.StreamID) error { + if _, ok := m.streams[id]; !ok { + return fmt.Errorf("tried to delete unknown incoming stream %d", id) + } + + // Don't delete this stream yet, if it was not yet accepted. + // Just save it to streamsToDelete map, to make sure it is deleted as soon as it gets accepted. + if id >= m.nextStreamToAccept { + entry, ok := m.streams[id] + if ok && entry.shouldDelete { + return fmt.Errorf("tried to delete incoming stream %d multiple times", id) + } + entry.shouldDelete = true + m.streams[id] = entry // can't assign to struct in map, so we need to reassign + return nil + } + + delete(m.streams, id) + // queue a MAX_STREAM_ID frame, giving the peer the option to open a new stream + if m.maxNumStreams > uint64(len(m.streams)) { + maxStream := m.nextStreamToOpen + 4*protocol.StreamID(m.maxNumStreams-uint64(len(m.streams))-1) + // never send a value larger than the maximum value for a stream number + if maxStream <= protocol.MaxStreamID { + m.maxStream = maxStream + m.queueMaxStreamID(&wire.MaxStreamsFrame{ + Type: m.streamType, + MaxStreamNum: m.maxStream.StreamNum(), + }) + } + } + return nil +} + +func (m *incomingStreamsMap[T]) CloseWithError(err error) { + m.mutex.Lock() + m.closeErr = err + for _, entry := range m.streams { + entry.stream.closeForShutdown(err) + } + m.mutex.Unlock() + close(m.newStreamChan) +} diff --git a/third_party/quic-go/streams_map_incoming_test.go b/third_party/quic-go/streams_map_incoming_test.go new file mode 100644 index 0000000..95386db --- /dev/null +++ b/third_party/quic-go/streams_map_incoming_test.go @@ -0,0 +1,364 @@ +package quic + +import ( + "context" + "math/rand/v2" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type mockStream struct { + id protocol.StreamID + + closed bool + closeErr error + sendWindow protocol.ByteCount + supportsResetStreamAt bool +} + +func (s *mockStream) closeForShutdown(err error) { + s.closed = true + s.closeErr = err +} + +func (s *mockStream) updateSendWindow(limit protocol.ByteCount) { + s.sendWindow = limit +} + +func (s *mockStream) enableResetStreamAt() { + s.supportsResetStreamAt = true +} + +func TestStreamsMapIncomingGettingStreams(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapIncomingGettingStreams(t, protocol.PerspectiveClient, protocol.FirstIncomingUniStreamClient) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapIncomingGettingStreams(t, protocol.PerspectiveServer, protocol.FirstIncomingUniStreamServer) + }) +} + +func testStreamsMapIncomingGettingStreams(t *testing.T, perspective protocol.Perspective, firstStream protocol.StreamID) { + var newStreamCounter int + const maxNumStreams = 10 + m := newIncomingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { + newStreamCounter++ + return &mockStream{id: id} + }, + maxNumStreams, + func(f wire.Frame) {}, + perspective, + ) + + // all streams up to the id on GetOrOpenStream are opened + str, err := m.GetOrOpenStream(firstStream + 4) + require.NoError(t, err) + require.NotNil(t, str) + require.Equal(t, 2, newStreamCounter) + require.Equal(t, firstStream+4, str.id) + // accept one of the streams + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err = m.AcceptStream(ctx) + require.NoError(t, err) + require.Equal(t, firstStream, str.id) + // open some more streams + str, err = m.GetOrOpenStream(firstStream + 16) + require.NoError(t, err) + require.Equal(t, 5, newStreamCounter) + require.Equal(t, firstStream+16, str.id) + // and accept all of them + for i := 1; i < 5; i++ { + str, err := m.AcceptStream(ctx) + require.NoError(t, err) + require.Equal(t, firstStream+4*protocol.StreamID(i), str.id) + } + + _, err = m.GetOrOpenStream(firstStream + 4*maxNumStreams - 4) + require.NoError(t, err) + _, err = m.GetOrOpenStream(firstStream + 4*maxNumStreams) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamLimitError}) + require.ErrorContains(t, err, "peer tried to open stream") + require.Equal(t, maxNumStreams, newStreamCounter) +} + +func TestStreamsMapIncomingAcceptingStreams(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + m := newIncomingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + 5, + func(f wire.Frame) {}, + protocol.PerspectiveClient, + ) + + // AcceptStream should respect the context + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + errChan := make(chan error, 1) + go func() { + _, err := m.AcceptStream(ctx) + errChan <- err + }() + + synctest.Wait() + + select { + case <-errChan: + t.Fatal("AcceptStream should not return") + default: + } + + cancel() + synctest.Wait() + select { + case err := <-errChan: + require.Equal(t, context.Canceled, err) + default: + t.Fatal("timeout") + } + + // AcceptStream should block if there are no streams available + go func() { + _, err := m.AcceptStream(context.Background()) + errChan <- err + }() + + synctest.Wait() + + select { + case <-errChan: + t.Fatal("AcceptStream should block") + default: + } + + _, err := m.GetOrOpenStream(protocol.FirstIncomingUniStreamClient) + require.NoError(t, err) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("timeout") + } + }) +} + +func TestStreamsMapIncomingDeletingStreams(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapIncomingDeletingStreams(t, protocol.PerspectiveClient, protocol.FirstIncomingUniStreamClient) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapIncomingDeletingStreams(t, protocol.PerspectiveServer, protocol.FirstIncomingUniStreamServer) + }) +} + +func testStreamsMapIncomingDeletingStreams(t *testing.T, perspective protocol.Perspective, firstStream protocol.StreamID) { + var frameQueue []wire.Frame + m := newIncomingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + 5, + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + perspective, + ) + err := m.DeleteStream(firstStream + 1337*4) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, "tried to delete unknown incoming stream") + + s, err := m.GetOrOpenStream(firstStream + 4) + require.NoError(t, err) + require.NotNil(t, s) + // delete the stream + require.NoError(t, m.DeleteStream(firstStream+4)) + require.Empty(t, frameQueue) + // it's not returned by GetOrOpenStream anymore + s, err = m.GetOrOpenStream(firstStream + 4) + require.NoError(t, err) + require.Nil(t, s) + + // AcceptStream still returns this stream + str, err := m.AcceptStream(context.Background()) + require.NoError(t, err) + require.Equal(t, firstStream, str.id) + require.Empty(t, frameQueue) + + str, err = m.AcceptStream(context.Background()) + require.NoError(t, err) + require.Equal(t, firstStream+4, str.id) + // now the stream is deleted and new stream credit is issued + require.Len(t, frameQueue, 1) + require.Equal(t, &wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: 6}, frameQueue[0]) + frameQueue = frameQueue[:0] + + require.NoError(t, m.DeleteStream(firstStream)) + require.Len(t, frameQueue, 1) + require.Equal(t, &wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: 7}, frameQueue[0]) +} + +// There's a maximum number that can be encoded in a MAX_STREAMS frame. +// Since the stream limit is configurable by the user, we can't rely on this number +// being high enough that it will never be reached in practice. +func TestStreamsMapIncomingDeletingStreamsWithHighLimits(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapIncomingDeletingStreamsWithHighLimits(t, protocol.PerspectiveClient, protocol.FirstIncomingUniStreamClient) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapIncomingDeletingStreamsWithHighLimits(t, protocol.PerspectiveServer, protocol.FirstIncomingUniStreamServer) + }) +} + +func testStreamsMapIncomingDeletingStreamsWithHighLimits(t *testing.T, pers protocol.Perspective, firstStream protocol.StreamID) { + var frameQueue []wire.Frame + m := newIncomingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + uint64(protocol.MaxStreamCount-2), + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + pers, + ) + + // open a bunch of streams + _, err := m.GetOrOpenStream(firstStream + 16) + require.NoError(t, err) + // accept all streams + for range 5 { + _, err := m.AcceptStream(context.Background()) + require.NoError(t, err) + } + require.Empty(t, frameQueue) + require.NoError(t, m.DeleteStream(firstStream+12)) + require.Len(t, frameQueue, 1) + require.Equal(t, + &wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: protocol.MaxStreamCount - 1}, + frameQueue[0], + ) + require.NoError(t, m.DeleteStream(firstStream+8)) + require.Len(t, frameQueue, 2) + require.Equal(t, + &wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: protocol.MaxStreamCount}, + frameQueue[1], + ) + // at this point, we can't increase the stream limit any further, so no more MAX_STREAMS frames will be sent + require.NoError(t, m.DeleteStream(firstStream+4)) + require.NoError(t, m.DeleteStream(firstStream)) + require.Len(t, frameQueue, 2) +} + +func TestStreamsMapIncomingClosing(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + m := newIncomingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + 5, + func(f wire.Frame) {}, + protocol.PerspectiveServer, + ) + + var streams []*mockStream + _, err := m.GetOrOpenStream(protocol.FirstIncomingUniStreamServer + 8) + require.NoError(t, err) + for range 3 { + str, err := m.AcceptStream(context.Background()) + require.NoError(t, err) + streams = append(streams, str) + } + + errChan := make(chan error, 1) + go func() { + _, err := m.AcceptStream(context.Background()) + errChan <- err + }() + + m.CloseWithError(assert.AnError) + synctest.Wait() + + // accepted streams should be closed + for _, str := range streams { + require.True(t, str.closed) + require.ErrorIs(t, str.closeErr, assert.AnError) + } + // AcceptStream should return the error + select { + case err := <-errChan: + require.ErrorIs(t, err, assert.AnError) + default: + t.Fatal("timeout") + } + }) +} + +func TestStreamsMapIncomingRandomized(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const num = 1000 + + streamType := []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi}[rand.IntN(2)] + firstStream := protocol.FirstIncomingUniStreamServer + if streamType == protocol.StreamTypeBidi { + firstStream = protocol.FirstIncomingBidiStreamServer + } + + m := newIncomingStreamsMap( + streamType, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + num, + func(f wire.Frame) {}, + protocol.PerspectiveServer, + ) + + ids := make([]protocol.StreamID, num) + for i := range num { + ids[i] = firstStream + 4*protocol.StreamID(i) + } + rand.Shuffle(len(ids), func(i, j int) { ids[i], ids[j] = ids[j], ids[i] }) + + errChan1 := make(chan error, 1) + go func() { + for range num { + if _, err := m.AcceptStream(context.Background()); err != nil { + errChan1 <- err + return + } + } + close(errChan1) + }() + + errChan2 := make(chan error, 1) + go func() { + for i := range num { + if _, err := m.GetOrOpenStream(ids[i]); err != nil { + errChan2 <- err + return + } + } + close(errChan2) + }() + + synctest.Wait() + + select { + case err := <-errChan1: + require.NoError(t, err) + default: + t.Fatal("should have accepted all streams") + } + select { + case err := <-errChan2: + require.NoError(t, err) + default: + t.Fatal("should have opened all streams") + } + }) +} diff --git a/third_party/quic-go/streams_map_outgoing.go b/third_party/quic-go/streams_map_outgoing.go new file mode 100644 index 0000000..b1a62af --- /dev/null +++ b/third_party/quic-go/streams_map_outgoing.go @@ -0,0 +1,253 @@ +package quic + +import ( + "context" + "fmt" + "slices" + "sync" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" +) + +type outgoingStream interface { + updateSendWindow(protocol.ByteCount) + enableResetStreamAt() + closeForShutdown(error) +} + +type outgoingStreamsMap[T outgoingStream] struct { + mutex sync.RWMutex + + streamType protocol.StreamType + streams map[protocol.StreamID]T + + openQueue []chan struct{} + + nextStream protocol.StreamID // stream ID of the stream returned by OpenStream(Sync) + maxStream protocol.StreamID // the maximum stream ID we're allowed to open + blockedSent bool // was a STREAMS_BLOCKED sent for the current maxStream + + newStream func(protocol.StreamID) T + queueStreamIDBlocked func(*wire.StreamsBlockedFrame) + + closeErr error +} + +func newOutgoingStreamsMap[T outgoingStream]( + streamType protocol.StreamType, + newStream func(protocol.StreamID) T, + queueControlFrame func(wire.Frame), + pers protocol.Perspective, +) *outgoingStreamsMap[T] { + var nextStream protocol.StreamID + switch { + case streamType == protocol.StreamTypeBidi && pers == protocol.PerspectiveServer: + nextStream = protocol.FirstOutgoingBidiStreamServer + case streamType == protocol.StreamTypeBidi && pers == protocol.PerspectiveClient: + nextStream = protocol.FirstOutgoingBidiStreamClient + case streamType == protocol.StreamTypeUni && pers == protocol.PerspectiveServer: + nextStream = protocol.FirstOutgoingUniStreamServer + case streamType == protocol.StreamTypeUni && pers == protocol.PerspectiveClient: + nextStream = protocol.FirstOutgoingUniStreamClient + } + return &outgoingStreamsMap[T]{ + streamType: streamType, + streams: make(map[protocol.StreamID]T), + maxStream: protocol.InvalidStreamNum, + nextStream: nextStream, + newStream: newStream, + queueStreamIDBlocked: func(f *wire.StreamsBlockedFrame) { queueControlFrame(f) }, + } +} + +func (m *outgoingStreamsMap[T]) OpenStream() (T, error) { + m.mutex.Lock() + defer m.mutex.Unlock() + + if m.closeErr != nil { + return *new(T), m.closeErr + } + + // if there are OpenStreamSync calls waiting, return an error here + if len(m.openQueue) > 0 || m.nextStream > m.maxStream { + m.maybeSendBlockedFrame() + return *new(T), &StreamLimitReachedError{} + } + return m.openStream(), nil +} + +func (m *outgoingStreamsMap[T]) OpenStreamSync(ctx context.Context) (T, error) { + m.mutex.Lock() + defer m.mutex.Unlock() + + if m.closeErr != nil { + return *new(T), m.closeErr + } + if err := ctx.Err(); err != nil { + return *new(T), err + } + if len(m.openQueue) == 0 && m.nextStream <= m.maxStream { + return m.openStream(), nil + } + + waitChan := make(chan struct{}, 1) + m.openQueue = append(m.openQueue, waitChan) + m.maybeSendBlockedFrame() + + for { + m.mutex.Unlock() + select { + case <-ctx.Done(): + m.mutex.Lock() + m.openQueue = slices.DeleteFunc(m.openQueue, func(c chan struct{}) bool { + return c == waitChan + }) + // If we just received a MAX_STREAMS frame, this might have been the next stream + // that could be opened. Make sure we unblock the next OpenStreamSync call. + m.maybeUnblockOpenSync() + return *new(T), ctx.Err() + case <-waitChan: + } + + m.mutex.Lock() + if m.closeErr != nil { + return *new(T), m.closeErr + } + if err := ctx.Err(); err != nil { + m.openQueue = slices.DeleteFunc(m.openQueue, func(c chan struct{}) bool { + return c == waitChan + }) + m.maybeUnblockOpenSync() + return *new(T), err + } + if m.nextStream > m.maxStream { + // no stream available. Continue waiting + continue + } + str := m.openStream() + m.openQueue = m.openQueue[1:] + m.maybeUnblockOpenSync() + return str, nil + } +} + +func (m *outgoingStreamsMap[T]) openStream() T { + s := m.newStream(m.nextStream) + m.streams[m.nextStream] = s + m.nextStream += 4 + return s +} + +// maybeSendBlockedFrame queues a STREAMS_BLOCKED frame for the current stream offset, +// if we haven't sent one for this offset yet +func (m *outgoingStreamsMap[T]) maybeSendBlockedFrame() { + if m.blockedSent { + return + } + + var streamLimit protocol.StreamNum + if m.maxStream != protocol.InvalidStreamID { + streamLimit = m.maxStream.StreamNum() + } + m.queueStreamIDBlocked(&wire.StreamsBlockedFrame{ + Type: m.streamType, + StreamLimit: streamLimit, + }) + m.blockedSent = true +} + +func (m *outgoingStreamsMap[T]) GetStream(id protocol.StreamID) (T, error) { + m.mutex.RLock() + if id >= m.nextStream { + m.mutex.RUnlock() + return *new(T), &qerr.TransportError{ + ErrorCode: qerr.StreamStateError, + ErrorMessage: fmt.Sprintf("peer attempted to open stream %d", id), + } + } + s := m.streams[id] + m.mutex.RUnlock() + return s, nil +} + +func (m *outgoingStreamsMap[T]) DeleteStream(id protocol.StreamID) error { + m.mutex.Lock() + defer m.mutex.Unlock() + + if _, ok := m.streams[id]; !ok { + return &qerr.TransportError{ + ErrorCode: qerr.StreamStateError, + ErrorMessage: fmt.Sprintf("tried to delete unknown outgoing stream %d", id), + } + } + delete(m.streams, id) + return nil +} + +func (m *outgoingStreamsMap[T]) SetMaxStream(id protocol.StreamID) { + m.mutex.Lock() + defer m.mutex.Unlock() + + if id <= m.maxStream { + return + } + m.maxStream = id + m.blockedSent = false + if m.maxStream < m.nextStream-4+4*protocol.StreamID(len(m.openQueue)) { + m.maybeSendBlockedFrame() + } + m.maybeUnblockOpenSync() +} + +// UpdateSendWindow is called when the peer's transport parameters are received. +// Only in the case of a 0-RTT handshake will we have open streams at this point. +// We might need to update the send window, in case the server increased it. +func (m *outgoingStreamsMap[T]) UpdateSendWindow(limit protocol.ByteCount) { + m.mutex.Lock() + for _, str := range m.streams { + str.updateSendWindow(limit) + } + m.mutex.Unlock() +} + +func (m *outgoingStreamsMap[T]) EnableResetStreamAt() { + m.mutex.Lock() + for _, str := range m.streams { + str.enableResetStreamAt() + } + m.mutex.Unlock() +} + +// unblockOpenSync unblocks the next OpenStreamSync go-routine to open a new stream +func (m *outgoingStreamsMap[T]) maybeUnblockOpenSync() { + if len(m.openQueue) == 0 { + return + } + if m.nextStream > m.maxStream { + return + } + // unblockOpenSync is called both from OpenStreamSync and from SetMaxStream. + // It's sufficient to only unblock OpenStreamSync once. + select { + case m.openQueue[0] <- struct{}{}: + default: + } +} + +func (m *outgoingStreamsMap[T]) CloseWithError(err error) { + m.mutex.Lock() + defer m.mutex.Unlock() + + m.closeErr = err + for _, str := range m.streams { + str.closeForShutdown(err) + } + for _, c := range m.openQueue { + if c != nil { + close(c) + } + } + m.openQueue = nil +} diff --git a/third_party/quic-go/streams_map_outgoing_test.go b/third_party/quic-go/streams_map_outgoing_test.go new file mode 100644 index 0000000..1b214ab --- /dev/null +++ b/third_party/quic-go/streams_map_outgoing_test.go @@ -0,0 +1,621 @@ +package quic + +import ( + "context" + "errors" + "fmt" + "math/rand/v2" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestStreamsMapOutgoingOpenAndDelete(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapOutgoingOpenAndDelete(t, protocol.PerspectiveClient, protocol.FirstOutgoingBidiStreamClient) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapOutgoingOpenAndDelete(t, protocol.PerspectiveServer, protocol.FirstOutgoingBidiStreamServer) + }) +} + +func testStreamsMapOutgoingOpenAndDelete(t *testing.T, perspective protocol.Perspective, firstStream protocol.StreamID) { + m := newOutgoingStreamsMap( + protocol.StreamTypeBidi, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) {}, + perspective, + ) + m.SetMaxStream(protocol.MaxStreamID) + + _, err := m.GetStream(firstStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("peer attempted to open stream %d", firstStream)) + + str1, err := m.OpenStream() + require.NoError(t, err) + require.Equal(t, firstStream, str1.id) + s, err := m.GetStream(firstStream) + require.NoError(t, err) + require.Equal(t, s, str1) + + str2, err := m.OpenStream() + require.NoError(t, err) + require.Equal(t, firstStream+4, str2.id) + + // update send window + m.UpdateSendWindow(1000) + require.Equal(t, protocol.ByteCount(1000), str1.sendWindow) + require.Equal(t, protocol.ByteCount(1000), str2.sendWindow) + + // enable reset stream at + m.EnableResetStreamAt() + require.True(t, str1.supportsResetStreamAt) + require.True(t, str2.supportsResetStreamAt) + + err = m.DeleteStream(firstStream + 1337*4) + require.Error(t, err) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, "tried to delete unknown outgoing stream") + + require.NoError(t, m.DeleteStream(firstStream)) + // deleting the same stream twice will fail + err = m.DeleteStream(firstStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, "tried to delete unknown outgoing stream") + // after deleting the stream it's not available anymore + str, err := m.GetStream(firstStream) + require.NoError(t, err) + require.Nil(t, str) +} + +func TestStreamsMapOutgoingLimits(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapOutgoingLimits(t, protocol.PerspectiveClient, protocol.FirstOutgoingUniStreamClient) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapOutgoingLimits(t, protocol.PerspectiveServer, protocol.FirstOutgoingUniStreamServer) + }) +} + +func testStreamsMapOutgoingLimits(t *testing.T, perspective protocol.Perspective, firstStream protocol.StreamID) { + synctest.Test(t, func(t *testing.T) { + m := newOutgoingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) {}, + perspective, + ) + m.SetMaxStream(firstStream) + + str, err := m.OpenStream() + require.NoError(t, err) + require.Equal(t, firstStream, str.id) + + // We've now reached the limit. OpenStream returns an error + _, err = m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + + // OpenStreamSync with a canceled context will return an error immediately + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, err = m.OpenStreamSync(ctx) + require.ErrorIs(t, err, context.Canceled) + + // OpenStreamSync blocks until the context is canceled... + ctx, cancel = context.WithCancel(context.Background()) + errChan := make(chan error, 1) + go func() { + _, err := m.OpenStreamSync(ctx) + errChan <- err + }() + + synctest.Wait() + select { + case <-errChan: + t.Fatal("didn't expect OpenStreamSync to return") + default: + } + // OpenStream still returns an error + _, err = m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + // cancelling the context unblocks OpenStreamSync + cancel() + synctest.Wait() + select { + case err := <-errChan: + require.ErrorIs(t, err, context.Canceled) + default: + t.Fatal("OpenStreamSync did not return after the context was canceled") + } + + // ... or until it's possible to open a new stream + var openedStream *mockStream + go func() { + str, err := m.OpenStreamSync(context.Background()) + openedStream = str + errChan <- err + }() + m.SetMaxStream(firstStream + 4) + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + require.Equal(t, firstStream+4, openedStream.id) + default: + t.Fatal("OpenStreamSync did not return after the stream limit was increased") + } + }) +} + +// This test checks that OpenStreamSync returns the context error when the context is canceled +// at the same time that the stream limit is increased (see https://github.com/apernet/quic-go/issues/5659). +// The race is inherently hard to trigger: even without the fix, this test only fails intermittently. +// To gain confidence in the fix, run it many times (e.g. 10000 times) with the race detector enabled. +func TestStreamsMapOutgoingOpenStreamSyncCancel(t *testing.T) { + queued := make(chan struct{}, 1) + m := newOutgoingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) { queued <- struct{}{} }, + protocol.PerspectiveClient, + ) + + ctx, cancel := context.WithCancel(context.Background()) + type result struct { + str *mockStream + err error + } + resultChan := make(chan result, 1) + go func() { + str, err := m.OpenStreamSync(ctx) + resultChan <- result{str: str, err: err} + }() + + select { + case <-queued: + case <-time.After(time.Second): + t.Fatal("OpenStreamSync did not queue") + } + + cancel() + m.SetMaxStream(protocol.FirstOutgoingUniStreamClient) + + select { + case res := <-resultChan: + require.ErrorIs(t, res.err, context.Canceled) + require.Nil(t, res.str) + case <-time.After(time.Second): + t.Fatal("OpenStreamSync did not return after the context was canceled") + } + + str, err := m.OpenStream() + require.NoError(t, err) + require.Equal(t, protocol.FirstOutgoingUniStreamClient, str.id) +} + +func TestStreamsMapOutgoingConcurrentOpenStreamSync(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + m := newOutgoingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) {}, + protocol.PerspectiveClient, + ) + + type result struct { + index int + stream *mockStream + err error + } + results := make(chan result, 3) + for i := range 3 { + go func(i int) { + str, err := m.OpenStreamSync(context.Background()) + results <- result{index: i, stream: str, err: err} + }(i) + time.Sleep(time.Minute) + } + + m.SetMaxStream(protocol.FirstOutgoingUniStreamClient + 4) + synctest.Wait() + received := make(map[protocol.StreamID]struct{}) + for range 2 { + select { + case res := <-results: + require.NoError(t, res.err) + require.Equal(t, protocol.FirstOutgoingUniStreamClient+4*protocol.StreamID(res.index), res.stream.id) + received[res.stream.id] = struct{}{} + default: + t.Fatal("OpenStreamSync did not return after the stream limit was increased") + } + } + require.Contains(t, received, protocol.FirstOutgoingUniStreamClient) + require.Contains(t, received, protocol.FirstOutgoingUniStreamClient+4) + + // the call to stream 3 is still blocked + select { + case <-results: + t.Fatal("expected OpenStreamSync to be blocked") + default: + } + m.SetMaxStream(protocol.FirstOutgoingUniStreamClient + 8) + synctest.Wait() + select { + case res := <-results: + require.NoError(t, res.err) + require.Equal(t, protocol.FirstOutgoingUniStreamClient+8, res.stream.id) + default: + t.Fatal("OpenStreamSync did not return after the stream limit was increased") + } + }) +} + +func TestStreamsMapOutgoingClosing(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + m := newOutgoingStreamsMap( + protocol.StreamTypeUni, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) {}, + protocol.PerspectiveServer, + ) + + m.SetMaxStream(protocol.FirstOutgoingUniStreamServer + 4) + str1, err := m.OpenStream() + require.NoError(t, err) + str2, err := m.OpenStream() + require.NoError(t, err) + + errChan := make(chan error, 1) + go func() { + _, err := m.OpenStreamSync(context.Background()) + errChan <- err + }() + + m.CloseWithError(assert.AnError) + + synctest.Wait() + + // both stream should be closed + assert.True(t, str1.closed) + assert.Equal(t, assert.AnError, str1.closeErr) + assert.True(t, str2.closed) + assert.Equal(t, assert.AnError, str2.closeErr) + + select { + case err := <-errChan: + require.Error(t, err) + default: + t.Fatal("OpenStreamSync did not return after the stream was closed") + } + }) +} + +func TestStreamsMapOutgoingBlockedFrames(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var frameQueue []wire.Frame + m := newOutgoingStreamsMap( + protocol.StreamTypeBidi, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) { frameQueue = append(frameQueue, f) }, + protocol.PerspectiveClient, + ) + + m.SetMaxStream(protocol.FirstOutgoingBidiStreamClient + 8) + for range 3 { + _, err := m.OpenStream() + require.NoError(t, err) + } + require.Empty(t, frameQueue) + + _, err := m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + require.Equal(t, []wire.Frame{ + &wire.StreamsBlockedFrame{Type: protocol.StreamTypeBidi, StreamLimit: 3}, + }, frameQueue) + frameQueue = frameQueue[:0] + + // only a single STREAMS_BLOCKED frame is queued per offset + for range 5 { + _, err = m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + require.Empty(t, frameQueue) + } + + errChan := make(chan error, 3) + for range 3 { + go func() { + _, err := m.OpenStreamSync(context.Background()) + errChan <- err + }() + } + synctest.Wait() + + // allow 2 more streams + m.SetMaxStream(protocol.FirstOutgoingBidiStreamClient + 16) + synctest.Wait() + + for range 2 { + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("OpenStreamSync did not return after the stream limit was increased") + } + } + require.Equal(t, []wire.Frame{ + &wire.StreamsBlockedFrame{Type: protocol.StreamTypeBidi, StreamLimit: 5}, + }, frameQueue) + frameQueue = frameQueue[:0] + + // now accept the last stream + m.SetMaxStream(protocol.FirstOutgoingBidiStreamClient + 20) + synctest.Wait() + select { + case err := <-errChan: + require.NoError(t, err) + default: + t.Fatal("OpenStreamSync did not return after the stream limit was increased") + } + require.Empty(t, frameQueue) + }) +} + +func TestStreamsMapOutgoingRandomizedOpenStreamSync(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + streamType := []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi}[rand.IntN(2)] + firstStream := protocol.FirstOutgoingUniStreamServer + if streamType == protocol.StreamTypeBidi { + firstStream = protocol.FirstOutgoingBidiStreamServer + } + + const n = 100 + + frameQueue := make(chan wire.Frame, n) + m := newOutgoingStreamsMap( + streamType, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) { frameQueue <- f }, + protocol.PerspectiveServer, + ) + + type result struct { + id protocol.StreamID + err error + } + resultChan := make(chan result, n) + for range n { + go func() { + str, err := m.OpenStreamSync(context.Background()) + resultChan <- result{id: str.id, err: err} + }() + } + synctest.Wait() + + select { + case f := <-frameQueue: + require.IsType(t, &wire.StreamsBlockedFrame{}, f) + require.Zero(t, f.(*wire.StreamsBlockedFrame).StreamLimit) + default: + t.Fatal("timed out waiting for STREAMS_BLOCKED frame") + } + + limit := firstStream - 4 + var limits []protocol.StreamID + seen := make(map[protocol.StreamID]struct{}) + maxStream := firstStream + 4*(n-1) + for limit < maxStream { + add := 4 * protocol.StreamID(rand.IntN(n/5)+1) + limit += add + if limit <= maxStream { + limits = append(limits, limit) + } + t.Logf("setting stream limit to %d", limit) + m.SetMaxStream(limit) + synctest.Wait() + + loop: + for { + select { + case res := <-resultChan: + require.NoError(t, res.err) + require.NotContains(t, seen, res.id) + require.LessOrEqual(t, res.id, limit) + seen[res.id] = struct{}{} + if len(seen) == int(limit.StreamNum()) || len(seen) == n { + break loop + } + default: + t.Fatalf("timed out waiting for stream to open") + } + } + + str, err := m.OpenStream() + if limit <= maxStream { + require.ErrorIs(t, err, &StreamLimitReachedError{}) + } else { + require.NoError(t, err) + require.Equal(t, maxStream+4, str.id) + } + } + require.Len(t, seen, n) + + close(frameQueue) + var blockedAt []protocol.StreamID + for f := range frameQueue { + if l := f.(*wire.StreamsBlockedFrame).StreamLimit; l <= n { + blockedAt = append(blockedAt, l.StreamID(streamType, protocol.PerspectiveServer)) + } + } + require.Equal(t, limits, blockedAt) + }) +} + +func TestStreamsMapOutgoingRandomizedWithCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const n = 100 + + streamType := []protocol.StreamType{protocol.StreamTypeUni, protocol.StreamTypeBidi}[rand.IntN(2)] + firstStream := protocol.FirstOutgoingUniStreamClient + if streamType == protocol.StreamTypeBidi { + firstStream = protocol.FirstOutgoingBidiStreamClient + } + + frameQueue := make(chan wire.Frame, n) + m := newOutgoingStreamsMap( + streamType, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) { frameQueue <- f }, + protocol.PerspectiveClient, + ) + + type result struct { + str *mockStream + err error + } + + ctx, cancel := context.WithCancel(context.Background()) + resultChan := make(chan result, 10*n) + var count int + var numCancelled int + for count < n { + shouldCancel := rand.IntN(n)%5 == 0 + if shouldCancel { + numCancelled++ + } else { + count++ + } + go func() { + var str *mockStream + var err error + if shouldCancel { + str, err = m.OpenStreamSync(ctx) + } else { + str, err = m.OpenStreamSync(context.Background()) + } + resultChan <- result{str: str, err: err} + }() + } + + synctest.Wait() + + select { + case f := <-frameQueue: + require.IsType(t, &wire.StreamsBlockedFrame{}, f) + require.Zero(t, f.(*wire.StreamsBlockedFrame).StreamLimit) + default: + t.Fatal("timed out waiting for STREAMS_BLOCKED frame") + } + + synctest.Wait() + cancel() + + limit := firstStream - 4 + maxStream := firstStream + 4*(n-1) + var limits []protocol.StreamID + seen := make(map[protocol.StreamID]struct{}) + var numCancelledSeen int + for limit < maxStream { + add := 4 * protocol.StreamID(rand.IntN(n/5)+1) + limit += add + if limit < maxStream { + limits = append(limits, limit) + } + t.Logf("setting stream limit to %d", limit) + m.SetMaxStream(limit) + + expectedOpened := int((min(maxStream, limit)-firstStream)/4) + 1 + for len(seen) < expectedOpened { + select { + case res := <-resultChan: + if errors.Is(res.err, context.Canceled) { + numCancelledSeen++ + } else { + require.NoError(t, res.err) + require.NotContains(t, seen, res.str.id) + require.LessOrEqual(t, res.str.id, min(maxStream, limit)) + seen[res.str.id] = struct{}{} + } + case <-time.After(time.Second): + t.Fatalf("timed out waiting for stream to open") + } + } + } + require.Len(t, seen, n) + for numCancelledSeen < numCancelled { + select { + case res := <-resultChan: + require.ErrorIs(t, res.err, context.Canceled) + numCancelledSeen++ + case <-time.After(time.Second): + t.Fatalf("timed out waiting for stream opening to be canceled") + } + } + t.Logf("saw %d streams, %d cancelled", len(seen), numCancelledSeen) + require.Equal(t, numCancelled, numCancelledSeen) + + close(frameQueue) + var blockedAt []protocol.StreamID + for f := range frameQueue { + sbf := f.(*wire.StreamsBlockedFrame) + require.Equal(t, streamType, sbf.Type) + blockedAt = append(blockedAt, sbf.StreamLimit.StreamID(streamType, protocol.PerspectiveClient)) + } + require.Equal(t, limits, blockedAt) + }) +} + +func TestStreamsMapConcurrent(t *testing.T) { + for i := range 5 { + t.Run(fmt.Sprintf("iteration %d", i+1), func(t *testing.T) { + testStreamsMapConcurrent(t) + }) + } +} + +func testStreamsMapConcurrent(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + m := newOutgoingStreamsMap( + protocol.StreamTypeBidi, + func(id protocol.StreamID) *mockStream { return &mockStream{id: id} }, + func(f wire.Frame) {}, + protocol.PerspectiveClient, + ) + + const num = 100 + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + errChan := make(chan error, num) + for range num { + go func() { + _, err := m.OpenStreamSync(ctx) + errChan <- err + }() + } + + go m.CloseWithError(assert.AnError) + go cancel() + go m.SetMaxStream(protocol.FirstOutgoingBidiStreamClient + 4*num/2) + + synctest.Wait() + + for range num { + select { + case err := <-errChan: + if err != nil { + require.True(t, errors.Is(err, assert.AnError) || errors.Is(err, context.Canceled)) + } + default: + t.Fatal("OpenStreamSync should have returned") + } + } + }) +} diff --git a/third_party/quic-go/streams_map_test.go b/third_party/quic-go/streams_map_test.go new file mode 100644 index 0000000..ee26ff7 --- /dev/null +++ b/third_party/quic-go/streams_map_test.go @@ -0,0 +1,674 @@ +package quic + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/wire" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func TestStreamsMapCreatingStreams(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapCreatingStreams(t, protocol.PerspectiveClient, + protocol.FirstIncomingBidiStreamClient, + protocol.FirstOutgoingBidiStreamClient, + protocol.FirstIncomingUniStreamClient, + protocol.FirstOutgoingUniStreamClient, + ) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapCreatingStreams(t, protocol.PerspectiveServer, + protocol.FirstIncomingBidiStreamServer, + protocol.FirstOutgoingBidiStreamServer, + protocol.FirstIncomingUniStreamServer, + protocol.FirstOutgoingUniStreamServer, + ) + }) +} + +func testStreamsMapCreatingStreams(t *testing.T, + perspective protocol.Perspective, + firstIncomingBidiStream protocol.StreamID, + firstOutgoingBidiStream protocol.StreamID, + firstIncomingUniStream protocol.StreamID, + firstOutgoingUniStream protocol.StreamID, +) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + m := newStreamsMap( + context.Background(), + mockSender, + func(wire.Frame) {}, + newTestStreamFlowController, + 1, + 1, + perspective, + ) + m.HandleTransportParameters(&wire.TransportParameters{ + MaxBidiStreamNum: protocol.MaxStreamCount, + MaxUniStreamNum: protocol.MaxStreamCount, + }) + + // opening streams + str1, err := m.OpenStream() + require.NoError(t, err) + str2, err := m.OpenStream() + require.NoError(t, err) + ustr1, err := m.OpenUniStream() + require.NoError(t, err) + ustr2, err := m.OpenUniStream() + require.NoError(t, err) + + assert.Equal(t, str1.StreamID(), firstOutgoingBidiStream) + assert.Equal(t, str2.StreamID(), firstOutgoingBidiStream+4) + assert.Equal(t, ustr1.StreamID(), firstOutgoingUniStream) + assert.Equal(t, ustr2.StreamID(), firstOutgoingUniStream+4) + + // accepting streams is triggered by receiving a frame referencing this stream + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: firstIncomingBidiStream}, monotime.Now())) + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: firstIncomingUniStream}, monotime.Now())) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + str, err := m.AcceptStream(ctx) + require.NoError(t, err) + ustr, err := m.AcceptUniStream(ctx) + require.NoError(t, err) + + assert.Equal(t, str.StreamID(), firstIncomingBidiStream) + assert.Equal(t, ustr.StreamID(), firstIncomingUniStream) +} + +func TestStreamsMapDeletingStreams(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapDeletingStreams(t, protocol.PerspectiveClient, + protocol.FirstIncomingBidiStreamClient, + protocol.FirstOutgoingBidiStreamClient, + protocol.FirstIncomingUniStreamClient, + protocol.FirstOutgoingUniStreamClient, + ) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapDeletingStreams(t, protocol.PerspectiveServer, + protocol.FirstIncomingBidiStreamServer, + protocol.FirstOutgoingBidiStreamServer, + protocol.FirstIncomingUniStreamServer, + protocol.FirstOutgoingUniStreamServer, + ) + }) +} + +func testStreamsMapDeletingStreams(t *testing.T, + perspective protocol.Perspective, + firstIncomingBidiStream protocol.StreamID, + firstOutgoingBidiStream protocol.StreamID, + firstIncomingUniStream protocol.StreamID, + firstOutgoingUniStream protocol.StreamID, +) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + var frameQueue []wire.Frame + m := newStreamsMap( + context.Background(), + mockSender, + func(frame wire.Frame) { frameQueue = append(frameQueue, frame) }, + newTestStreamFlowController, + 100, + 100, + perspective, + ) + m.HandleTransportParameters(&wire.TransportParameters{ + MaxBidiStreamNum: 10, + MaxUniStreamNum: 10, + }) + + _, err := m.OpenStream() + require.NoError(t, err) + require.NoError(t, m.DeleteStream(firstOutgoingBidiStream)) + err = m.DeleteStream(firstOutgoingBidiStream + 400) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("tried to delete unknown outgoing stream %d", firstOutgoingBidiStream+400)) + + _, err = m.OpenUniStream() + require.NoError(t, err) + require.NoError(t, m.DeleteStream(firstOutgoingUniStream)) + err = m.DeleteStream(firstOutgoingUniStream + 400) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("tried to delete unknown outgoing stream %d", firstOutgoingUniStream+400)) + + require.Empty(t, frameQueue) + // deleting incoming bidirectional streams + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: firstIncomingBidiStream}, monotime.Now())) + require.NoError(t, m.DeleteStream(firstIncomingBidiStream)) + err = m.DeleteStream(firstIncomingBidiStream + 400) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("tried to delete unknown incoming stream %d", firstIncomingBidiStream+400)) + + // the MAX_STREAMS frame is only queued once the stream is accepted + require.Empty(t, frameQueue) + _, err = m.AcceptStream(context.Background()) + require.NoError(t, err) + + require.Equal(t, frameQueue, []wire.Frame{ + &wire.MaxStreamsFrame{ + Type: protocol.StreamTypeBidi, + MaxStreamNum: 101, + }, + }) + frameQueue = frameQueue[:0] + + // deleting incoming unidirectional streams + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: firstIncomingUniStream}, monotime.Now())) + require.NoError(t, m.DeleteStream(firstIncomingUniStream)) + err = m.DeleteStream(firstIncomingUniStream + 400) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("tried to delete unknown incoming stream %d", firstIncomingUniStream+400)) + + // the MAX_STREAMS frame is only queued once the stream is accepted + require.Empty(t, frameQueue) + _, err = m.AcceptUniStream(context.Background()) + require.NoError(t, err) + + require.Equal(t, frameQueue, []wire.Frame{ + &wire.MaxStreamsFrame{ + Type: protocol.StreamTypeUni, + MaxStreamNum: 101, + }, + }) + frameQueue = frameQueue[:0] +} + +func TestStreamsMapStreamLimits(t *testing.T) { + t.Run("client", func(t *testing.T) { + testStreamsMapStreamLimits(t, protocol.PerspectiveClient) + }) + t.Run("server", func(t *testing.T) { + testStreamsMapStreamLimits(t, protocol.PerspectiveServer) + }) +} + +func testStreamsMapStreamLimits(t *testing.T, perspective protocol.Perspective) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + var frameQueue []wire.Frame + m := newStreamsMap( + context.Background(), + mockSender, + func(frame wire.Frame) { frameQueue = append(frameQueue, frame) }, + newTestStreamFlowController, + 100, + 100, + perspective, + ) + + // increase via transport parameters + _, err := m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + require.ErrorContains(t, err, "too many open streams") + m.HandleTransportParameters(&wire.TransportParameters{MaxBidiStreamNum: 1}) + _, err = m.OpenStream() + require.NoError(t, err) + _, err = m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + + _, err = m.OpenUniStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + m.HandleTransportParameters(&wire.TransportParameters{MaxUniStreamNum: 1}) + _, err = m.OpenUniStream() + require.NoError(t, err) + _, err = m.OpenUniStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + + // increase via MAX_STREAMS frames + m.HandleMaxStreamsFrame(&wire.MaxStreamsFrame{ + Type: protocol.StreamTypeBidi, + MaxStreamNum: 2, + }) + _, err = m.OpenStream() + require.NoError(t, err) + _, err = m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + + m.HandleMaxStreamsFrame(&wire.MaxStreamsFrame{ + Type: protocol.StreamTypeUni, + MaxStreamNum: 2, + }) + _, err = m.OpenUniStream() + require.NoError(t, err) + _, err = m.OpenUniStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) + + // decrease via transport parameters + m.HandleTransportParameters(&wire.TransportParameters{MaxBidiStreamNum: 0}) + _, err = m.OpenStream() + require.ErrorIs(t, err, &StreamLimitReachedError{}) +} + +func TestStreamsMapHandleReceiveStreamFrames(t *testing.T) { + for _, pers := range []protocol.Perspective{protocol.PerspectiveClient, protocol.PerspectiveServer} { + t.Run(pers.String(), func(t *testing.T) { + t.Run("STREAM frame", func(t *testing.T) { + testStreamsMapHandleReceiveStreamFrames(t, + pers, + func(m *streamsMap, id protocol.StreamID) error { + return m.HandleStreamFrame(&wire.StreamFrame{StreamID: id}, monotime.Now()) + }, + ) + }) + + t.Run("STREAM_DATA_BLOCKED frame", func(t *testing.T) { + testStreamsMapHandleReceiveStreamFrames(t, + pers, + func(m *streamsMap, id protocol.StreamID) error { + return m.HandleStreamDataBlockedFrame(&wire.StreamDataBlockedFrame{StreamID: id}) + }, + ) + }) + + t.Run("RESET_STREAM frame", func(t *testing.T) { + testStreamsMapHandleReceiveStreamFrames(t, + pers, + func(m *streamsMap, id protocol.StreamID) error { + return m.HandleResetStreamFrame(&wire.ResetStreamFrame{StreamID: id}, monotime.Now()) + }, + ) + }) + }) + } +} + +func testStreamsMapHandleReceiveStreamFrames(t *testing.T, pers protocol.Perspective, handleFrame func(*streamsMap, protocol.StreamID) error) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + var streamsCreated []protocol.StreamID + m := newStreamsMap( + context.Background(), + mockSender, + func(frame wire.Frame) {}, + func(id protocol.StreamID) *streamFlowController { + streamsCreated = append(streamsCreated, id) + return newTestStreamFlowController(id) + }, + 100, + 100, + pers, + ) + m.HandleMaxStreamsFrame(&wire.MaxStreamsFrame{Type: protocol.StreamTypeBidi, MaxStreamNum: protocol.MaxStreamCount}) + m.HandleMaxStreamsFrame(&wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: protocol.MaxStreamCount}) + + var firstOutgoingUniStream, firstOutgoingBidiStream, firstIncomingUniStream, firstIncomingBidiStream protocol.StreamID + if pers == protocol.PerspectiveClient { + firstOutgoingBidiStream = protocol.FirstOutgoingBidiStreamClient + firstOutgoingUniStream = protocol.FirstOutgoingUniStreamClient + firstIncomingUniStream = protocol.FirstIncomingUniStreamClient + firstIncomingBidiStream = protocol.FirstIncomingBidiStreamClient + } else { + firstOutgoingBidiStream = protocol.FirstOutgoingBidiStreamServer + firstOutgoingUniStream = protocol.FirstOutgoingUniStreamServer + firstIncomingUniStream = protocol.FirstIncomingUniStreamServer + firstIncomingBidiStream = protocol.FirstIncomingBidiStreamServer + } + + // 1. The peer can't open a unidirectional send stream... + err := handleFrame(m, firstOutgoingUniStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("invalid frame for receive stream %d", firstOutgoingUniStream)) + require.Empty(t, streamsCreated) + // ... and a STREAM frame for a unidirectional send stream is invalid even if the stream is open. + _, err = m.OpenUniStream() + require.NoError(t, err) + err = handleFrame(m, firstOutgoingUniStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("invalid frame for receive stream %d", firstOutgoingUniStream)) + streamsCreated = streamsCreated[:0] + + // 2. The peer can't open a bidirectional stream initiated by us... + err = handleFrame(m, firstOutgoingBidiStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("peer attempted to open stream %d", firstOutgoingBidiStream)) + require.Empty(t, streamsCreated) + // ... but it's valid once we have opened the stream. + _, err = m.OpenStream() + require.NoError(t, err) + require.NoError(t, handleFrame(m, firstOutgoingBidiStream)) + streamsCreated = streamsCreated[:0] + // Delayed frames for deleted streams are absorbed. + require.NoError(t, m.DeleteStream(firstOutgoingBidiStream)) + require.NoError(t, handleFrame(m, firstOutgoingBidiStream)) + require.Empty(t, streamsCreated) + + // 3. The peer can send STREAM frames for unidirectional receive streams, + // as long as they're below the stream limit. + require.ErrorIs(t, + handleFrame(m, firstIncomingUniStream+400), + &qerr.TransportError{ErrorCode: qerr.StreamLimitError}, + ) + require.Empty(t, streamsCreated) + require.NoError(t, handleFrame(m, firstIncomingUniStream)) + require.Equal(t, streamsCreated, []protocol.StreamID{firstIncomingUniStream}) + streamsCreated = streamsCreated[:0] + // Delayed frames for deleted streams are absorbed. + require.NoError(t, m.DeleteStream(firstIncomingUniStream)) + require.NoError(t, handleFrame(m, firstIncomingUniStream)) + require.Empty(t, streamsCreated) + + // 4. The peer can send STREAM frames for bidirectional receive streams, + // as long as they're below the stream limit. + require.ErrorIs(t, + handleFrame(m, firstIncomingBidiStream+400), + &qerr.TransportError{ErrorCode: qerr.StreamLimitError}, + ) + require.Empty(t, streamsCreated) + require.NoError(t, handleFrame(m, firstIncomingBidiStream)) + require.Equal(t, streamsCreated, []protocol.StreamID{firstIncomingBidiStream}) +} + +func TestStreamsMapHandleSendStreamFrames(t *testing.T) { + for _, pers := range []protocol.Perspective{protocol.PerspectiveClient, protocol.PerspectiveServer} { + t.Run(pers.String(), func(t *testing.T) { + t.Run("STOP_SENDING frame", func(t *testing.T) { + testStreamsMapHandleSendStreamFrames(t, + pers, + func(m *streamsMap, id protocol.StreamID) error { + return m.HandleStopSendingFrame(&wire.StopSendingFrame{StreamID: id}) + }, + ) + }) + + t.Run("MAX_STREAM_DATA frame", func(t *testing.T) { + testStreamsMapHandleSendStreamFrames(t, + pers, + func(m *streamsMap, id protocol.StreamID) error { + return m.HandleMaxStreamDataFrame(&wire.MaxStreamDataFrame{StreamID: id, MaximumStreamData: 1000}) + }, + ) + }) + }) + } +} + +func testStreamsMapHandleSendStreamFrames(t *testing.T, pers protocol.Perspective, handleFrame func(m *streamsMap, id protocol.StreamID) error) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + mockSender.EXPECT().onHasStreamControlFrame(gomock.Any(), gomock.Any()).AnyTimes() + var streamsCreated []protocol.StreamID + m := newStreamsMap( + context.Background(), + mockSender, + func(frame wire.Frame) {}, + func(id protocol.StreamID) *streamFlowController { + streamsCreated = append(streamsCreated, id) + return newTestStreamFlowController(id) + }, + 100, + 100, + pers, + ) + m.HandleMaxStreamsFrame(&wire.MaxStreamsFrame{Type: protocol.StreamTypeBidi, MaxStreamNum: protocol.MaxStreamCount}) + m.HandleMaxStreamsFrame(&wire.MaxStreamsFrame{Type: protocol.StreamTypeUni, MaxStreamNum: protocol.MaxStreamCount}) + + var firstOutgoingUniStream, firstOutgoingBidiStream, firstIncomingUniStream, firstIncomingBidiStream protocol.StreamID + if pers == protocol.PerspectiveClient { + firstOutgoingBidiStream = protocol.FirstOutgoingBidiStreamClient + firstOutgoingUniStream = protocol.FirstOutgoingUniStreamClient + firstIncomingUniStream = protocol.FirstIncomingUniStreamClient + firstIncomingBidiStream = protocol.FirstIncomingBidiStreamClient + } else { + firstOutgoingBidiStream = protocol.FirstOutgoingBidiStreamServer + firstOutgoingUniStream = protocol.FirstOutgoingUniStreamServer + firstIncomingUniStream = protocol.FirstIncomingUniStreamServer + firstIncomingBidiStream = protocol.FirstIncomingBidiStreamServer + } + + // 1. The peer can't open a unidirectional send stream... + err := handleFrame(m, firstOutgoingUniStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("peer attempted to open stream %d", firstOutgoingUniStream)) + require.Empty(t, streamsCreated) + // ... but once we have opened the stream, it's valid. + _, err = m.OpenUniStream() + require.NoError(t, err) + require.NoError(t, handleFrame(m, firstOutgoingUniStream)) + streamsCreated = streamsCreated[:0] + // Delayed frames for deleted streams are absorbed. + require.NoError(t, m.DeleteStream(firstOutgoingUniStream)) + require.NoError(t, handleFrame(m, firstOutgoingUniStream)) + require.Empty(t, streamsCreated) + + // 2. The peer can't open a bidirectional stream initiated by us... + err = handleFrame(m, firstOutgoingBidiStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("peer attempted to open stream %d", firstOutgoingBidiStream)) + require.Empty(t, streamsCreated) + // ... but once we have opened the stream, it's valid. + _, err = m.OpenStream() + require.NoError(t, err) + require.NoError(t, handleFrame(m, firstOutgoingBidiStream)) + streamsCreated = streamsCreated[:0] + // Delayed frames for deleted streams are absorbed. + require.NoError(t, m.DeleteStream(firstOutgoingBidiStream)) + require.NoError(t, handleFrame(m, firstOutgoingBidiStream)) + require.Empty(t, streamsCreated) + + // 3. The peer can't send STOP_SENDING frames for unidirectional send streams + err = handleFrame(m, firstIncomingUniStream) + require.ErrorIs(t, err, &qerr.TransportError{ErrorCode: qerr.StreamStateError}) + require.ErrorContains(t, err, fmt.Sprintf("invalid frame for send stream %d", firstIncomingUniStream)) + require.Empty(t, streamsCreated) + + // 4. The peer can send STOP_SENDING frames for bidirectional receive streams iniated by itself, + // as long as they're below the stream limit. + require.ErrorIs(t, + handleFrame(m, firstIncomingBidiStream+400), + &qerr.TransportError{ErrorCode: qerr.StreamLimitError}, + ) + require.Empty(t, streamsCreated) + require.NoError(t, handleFrame(m, firstIncomingBidiStream)) + require.Equal(t, streamsCreated, []protocol.StreamID{firstIncomingBidiStream}) + streamsCreated = streamsCreated[:0] + // Delayed frames for deleted streams are absorbed. + require.NoError(t, m.DeleteStream(firstIncomingBidiStream)) + require.NoError(t, handleFrame(m, firstIncomingBidiStream)) + require.Empty(t, streamsCreated) +} + +func TestStreamsMapClosing(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + m := newStreamsMap( + context.Background(), + mockSender, + func(wire.Frame) {}, + newTestStreamFlowController, + 1, + 1, + protocol.PerspectiveClient, + ) + m.CloseWithError(assert.AnError) + _, err := m.OpenStream() + require.ErrorIs(t, err, assert.AnError) + _, err = m.OpenUniStream() + require.ErrorIs(t, err, assert.AnError) + _, err = m.AcceptStream(context.Background()) + require.ErrorIs(t, err, assert.AnError) + _, err = m.AcceptUniStream(context.Background()) + require.ErrorIs(t, err, assert.AnError) +} + +func TestStreamsMap0RTT(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + var fcs []*streamFlowController + m := newStreamsMap( + context.Background(), + mockSender, + func(wire.Frame) {}, + func(id protocol.StreamID) *streamFlowController { + fc := newTestStreamFlowController(id) + fcs = append(fcs, fc) + return fc + }, + 1, + 1, + protocol.PerspectiveClient, + ) + // restored transport parameters + m.HandleTransportParameters(&wire.TransportParameters{ + MaxBidiStreamNum: 1, + MaxUniStreamNum: 1, + }) + _, err := m.OpenStream() + require.NoError(t, err) + _, err = m.OpenUniStream() + require.NoError(t, err) + + // new transport parameters + m.HandleTransportParameters(&wire.TransportParameters{ + MaxBidiStreamNum: 1000, + InitialMaxStreamDataBidiRemote: 1234, + MaxUniStreamNum: 1000, + InitialMaxStreamDataUni: 4321, + }) + require.Len(t, fcs, 2) + require.Equal(t, protocol.ByteCount(1234), fcs[0].SendWindowSize()) + require.Equal(t, protocol.ByteCount(4321), fcs[1].SendWindowSize()) +} + +func TestStreamsMap0RTTResetStreamAt(t *testing.T) { + for _, enabled := range []bool{false, true} { + t.Run(fmt.Sprintf("enabled: %t", enabled), func(t *testing.T) { + mockSender := NewMockStreamSender(gomock.NewController(t)) + mockSender.EXPECT().onHasStreamData(gomock.Any(), gomock.Any()).AnyTimes() + mockSender.EXPECT().onHasStreamControlFrame(gomock.Any(), gomock.Any()).AnyTimes() + m := newStreamsMap( + context.Background(), + mockSender, + func(wire.Frame) {}, + func(id protocol.StreamID) *streamFlowController { + return newTestStreamFlowControllerWithSendWindow(id, 1) + }, + 1, + 1, + protocol.PerspectiveClient, + ) + m.HandleTransportParameters(&wire.TransportParameters{MaxBidiStreamNum: 1, MaxUniStreamNum: 1}) + str, err := m.OpenStream() + require.NoError(t, err) + uniStr, err := m.OpenUniStream() + require.NoError(t, err) + + m.HandleTransportParameters(&wire.TransportParameters{EnableResetStreamAt: enabled}) + require.Equal(t, enabled, supportsResetStreamAt(t, str)) + require.Equal(t, enabled, sendStreamSupportsResetStreamAt(t, uniStr)) + }) + } +} + +func TestStreamsMap0RTTRejection(t *testing.T) { + mockCtrl := gomock.NewController(t) + mockSender := NewMockStreamSender(mockCtrl) + m := newStreamsMap( + context.Background(), + mockSender, + func(wire.Frame) {}, + newTestStreamFlowController, + 1, + 1, + protocol.PerspectiveClient, + ) + + m.ResetFor0RTT() + _, err := m.OpenStream() + require.ErrorIs(t, err, Err0RTTRejected) + _, err = m.OpenUniStream() + require.ErrorIs(t, err, Err0RTTRejected) + _, err = m.AcceptStream(context.Background()) + require.ErrorIs(t, err, Err0RTTRejected) + _, err = m.AcceptUniStream(context.Background()) + require.ErrorIs(t, err, Err0RTTRejected) + + // make sure that we can still get new streams, as the server might be sending us data + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: 3}, monotime.Now())) + + // now switch to using the new streams map + m.UseResetMaps() + _, err = m.OpenStream() + require.Error(t, err) + require.ErrorIs(t, err, &StreamLimitReachedError{}) +} + +func TestStreamsMap0RTTRejectionResetStreamAt(t *testing.T) { + for _, enabled := range []bool{false, true} { + t.Run(fmt.Sprintf("enabled: %t", enabled), func(t *testing.T) { + testStreamsMap0RTTRejectionResetStreamAt(t, enabled) + }) + } +} + +func testStreamsMap0RTTRejectionResetStreamAt(t *testing.T, enabled bool) { + mockSender := NewMockStreamSender(gomock.NewController(t)) + mockSender.EXPECT().onHasStreamData(gomock.Any(), gomock.Any()).AnyTimes() + mockSender.EXPECT().onHasStreamControlFrame(gomock.Any(), gomock.Any()).AnyTimes() + m := newStreamsMap( + context.Background(), + mockSender, + func(wire.Frame) {}, + func(id protocol.StreamID) *streamFlowController { + return newTestStreamFlowControllerWithSendWindow(id, 1) + }, + 2, + 1, + protocol.PerspectiveClient, + ) + m.HandleTransportParameters(&wire.TransportParameters{EnableResetStreamAt: true}) + m.ResetFor0RTT() + + // The server can send 0.5-RTT data before the handshake completes. + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: 1}, monotime.Now())) + m.UseResetMaps() + + str, err := m.AcceptStream(context.Background()) + require.NoError(t, err) + require.False(t, supportsResetStreamAt(t, str)) + + m.HandleTransportParameters(&wire.TransportParameters{EnableResetStreamAt: enabled}) + require.NoError(t, m.HandleStreamFrame(&wire.StreamFrame{StreamID: 5}, monotime.Now())) + str, err = m.AcceptStream(context.Background()) + require.NoError(t, err) + require.Equal(t, enabled, supportsResetStreamAt(t, str)) +} + +func supportsResetStreamAt(t *testing.T, str *Stream) bool { + t.Helper() + _, err := str.Write([]byte{0}) + require.NoError(t, err) + str.SetReliableBoundary() + str.CancelWrite(0) + frame, ok, _ := str.getControlFrame(monotime.Now()) + require.True(t, ok) + reset, ok := frame.Frame.(*wire.ResetStreamFrame) + require.True(t, ok) + return reset.ReliableSize > 0 +} + +func sendStreamSupportsResetStreamAt(t *testing.T, str *SendStream) bool { + t.Helper() + _, err := str.Write([]byte{0}) + require.NoError(t, err) + str.SetReliableBoundary() + str.CancelWrite(0) + frame, ok, _ := str.getControlFrame(monotime.Now()) + require.True(t, ok) + reset, ok := frame.Frame.(*wire.ResetStreamFrame) + require.True(t, ok) + return reset.ReliableSize > 0 +} diff --git a/third_party/quic-go/sys_conn.go b/third_party/quic-go/sys_conn.go new file mode 100644 index 0000000..806b1a2 --- /dev/null +++ b/third_party/quic-go/sys_conn.go @@ -0,0 +1,122 @@ +package quic + +import ( + "io" + "net" + "syscall" + "time" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +type connCapabilities struct { + // This connection has the Don't Fragment (DF) bit set. + // This means it makes to run DPLPMTUD. + DF bool + // GSO (Generic Segmentation Offload) supported + GSO bool + // ECN (Explicit Congestion Notifications) supported + ECN bool +} + +// rawConn is a connection that allow reading of a receivedPackeh. +type rawConn interface { + ReadPacket() (receivedPacket, error) + // WritePacket writes a packet on the wire. + // gsoSize is the size of a single packet, or 0 to disable GSO. + // It is invalid to set gsoSize if capabilities.GSO is not set. + WritePacket(b []byte, addr net.Addr, packetInfoOOB []byte, gsoSize uint16, ecn protocol.ECN) (int, error) + LocalAddr() net.Addr + SetReadDeadline(time.Time) error + io.Closer + + capabilities() connCapabilities +} + +// OOBCapablePacketConn is a connection that allows the reading of ECN bits from the IP header. +// If the PacketConn passed to the [Transport] satisfies this interface, quic-go will use it. +// In this case, ReadMsgUDP() will be used instead of ReadFrom() to read packets. +type OOBCapablePacketConn interface { + net.PacketConn + SyscallConn() (syscall.RawConn, error) + SetReadBuffer(int) error + ReadMsgUDP(b, oob []byte) (n, oobn, flags int, addr *net.UDPAddr, err error) + WriteMsgUDP(b, oob []byte, addr *net.UDPAddr) (n, oobn int, err error) +} + +var _ OOBCapablePacketConn = &net.UDPConn{} + +func wrapConn(pc net.PacketConn, disableGSO bool) (rawConn, error) { + _ = setReceiveBuffer(pc) + _ = setSendBuffer(pc) + + conn, ok := pc.(interface { + SyscallConn() (syscall.RawConn, error) + }) + var supportsDF bool = true + if ok { + rawConn, err := conn.SyscallConn() + if err != nil { + return nil, err + } + + // only set DF on UDP sockets + if _, ok := pc.LocalAddr().(*net.UDPAddr); ok { + var err error + supportsDF, err = setDF(rawConn) + if err != nil { + return nil, err + } + } + } + c, ok := pc.(OOBCapablePacketConn) + if !ok { + utils.DefaultLogger.Infof("PacketConn is not a net.UDPConn. Disabling optimizations possible on UDP connections.") + return &basicConn{PacketConn: pc, supportsDF: supportsDF}, nil + } + return newConn(c, supportsDF, disableGSO) +} + +// The basicConn is the most trivial implementation of a rawConn. +// It reads a single packet from the underlying net.PacketConn. +// It is used when +// * the net.PacketConn is not a OOBCapablePacketConn, and +// * when the OS doesn't support OOB. +type basicConn struct { + net.PacketConn + supportsDF bool +} + +var _ rawConn = &basicConn{} + +func (c *basicConn) ReadPacket() (receivedPacket, error) { + buffer := getPacketBuffer() + // The packet size should not exceed protocol.MaxPacketBufferSize bytes + // If it does, we only read a truncated packet, which will then end up undecryptable + buffer.Data = buffer.Data[:protocol.MaxPacketBufferSize] + n, addr, err := c.ReadFrom(buffer.Data) + if err != nil { + buffer.Release() + return receivedPacket{}, err + } + return receivedPacket{ + remoteAddr: addr, + rcvTime: monotime.Now(), + data: buffer.Data[:n], + buffer: buffer, + }, nil +} + +func (c *basicConn) WritePacket(b []byte, addr net.Addr, _ []byte, gsoSize uint16, ecn protocol.ECN) (n int, err error) { + if gsoSize != 0 { + panic("cannot use GSO with a basicConn") + } + if ecn != protocol.ECNUnsupported { + panic("cannot use ECN with a basicConn") + } + return c.WriteTo(b, addr) +} + +func (c *basicConn) capabilities() connCapabilities { return connCapabilities{DF: c.supportsDF} } diff --git a/third_party/quic-go/sys_conn_buffers.go b/third_party/quic-go/sys_conn_buffers.go new file mode 100644 index 0000000..b2b64fa --- /dev/null +++ b/third_party/quic-go/sys_conn_buffers.go @@ -0,0 +1,68 @@ +package quic + +import ( + "errors" + "fmt" + "net" + "syscall" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +//go:generate sh -c "echo '// Code generated by go generate. DO NOT EDIT.\n// Source: sys_conn_buffers.go\n' > sys_conn_buffers_write.go && sed -e 's/SetReadBuffer/SetWriteBuffer/g' -e 's/setReceiveBuffer/setSendBuffer/g' -e 's/inspectReadBuffer/inspectWriteBuffer/g' -e 's/protocol\\.DesiredReceiveBufferSize/protocol\\.DesiredSendBufferSize/g' -e 's/forceSetReceiveBuffer/forceSetSendBuffer/g' -e 's/receive buffer/send buffer/g' sys_conn_buffers.go | sed '/^\\/\\/go:generate/d' >> sys_conn_buffers_write.go" +func setReceiveBuffer(c net.PacketConn) error { + conn, ok := c.(interface{ SetReadBuffer(int) error }) + if !ok { + return errors.New("connection doesn't allow setting of receive buffer size. Not a *net.UDPConn?") + } + + var syscallConn syscall.RawConn + if sc, ok := c.(interface { + SyscallConn() (syscall.RawConn, error) + }); ok { + var err error + syscallConn, err = sc.SyscallConn() + if err != nil { + syscallConn = nil + } + } + // The connection has a SetReadBuffer method, but we couldn't obtain a syscall.RawConn. + // This shouldn't happen for a net.UDPConn, but is possible if the connection just implements the + // net.PacketConn interface and the SetReadBuffer method. + // We have no way of checking if increasing the buffer size actually worked. + if syscallConn == nil { + return conn.SetReadBuffer(protocol.DesiredReceiveBufferSize) + } + + size, err := inspectReadBuffer(syscallConn) + if err != nil { + return fmt.Errorf("failed to determine receive buffer size: %w", err) + } + if size >= protocol.DesiredReceiveBufferSize { + utils.DefaultLogger.Debugf("Conn has receive buffer of %d kiB (wanted: at least %d kiB)", size/1024, protocol.DesiredReceiveBufferSize/1024) + return nil + } + // Ignore the error. We check if we succeeded by querying the buffer size afterward. + _ = conn.SetReadBuffer(protocol.DesiredReceiveBufferSize) + newSize, err := inspectReadBuffer(syscallConn) + if newSize < protocol.DesiredReceiveBufferSize { + // Try again with RCVBUFFORCE on Linux + _ = forceSetReceiveBuffer(syscallConn, protocol.DesiredReceiveBufferSize) + newSize, err = inspectReadBuffer(syscallConn) + if err != nil { + return fmt.Errorf("failed to determine receive buffer size: %w", err) + } + } + if err != nil { + return fmt.Errorf("failed to determine receive buffer size: %w", err) + } + if newSize == size { + return fmt.Errorf("failed to increase receive buffer size (wanted: %d kiB, got %d kiB)", protocol.DesiredReceiveBufferSize/1024, newSize/1024) + } + if newSize < protocol.DesiredReceiveBufferSize { + return fmt.Errorf("failed to sufficiently increase receive buffer size (was: %d kiB, wanted: %d kiB, got: %d kiB)", size/1024, protocol.DesiredReceiveBufferSize/1024, newSize/1024) + } + utils.DefaultLogger.Debugf("Increased receive buffer size to %d kiB", newSize/1024) + return nil +} diff --git a/third_party/quic-go/sys_conn_buffers_write.go b/third_party/quic-go/sys_conn_buffers_write.go new file mode 100644 index 0000000..0468bbe --- /dev/null +++ b/third_party/quic-go/sys_conn_buffers_write.go @@ -0,0 +1,70 @@ +// Code generated by go generate. DO NOT EDIT. +// Source: sys_conn_buffers.go + +package quic + +import ( + "errors" + "fmt" + "net" + "syscall" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +func setSendBuffer(c net.PacketConn) error { + conn, ok := c.(interface{ SetWriteBuffer(int) error }) + if !ok { + return errors.New("connection doesn't allow setting of send buffer size. Not a *net.UDPConn?") + } + + var syscallConn syscall.RawConn + if sc, ok := c.(interface { + SyscallConn() (syscall.RawConn, error) + }); ok { + var err error + syscallConn, err = sc.SyscallConn() + if err != nil { + syscallConn = nil + } + } + // The connection has a SetWriteBuffer method, but we couldn't obtain a syscall.RawConn. + // This shouldn't happen for a net.UDPConn, but is possible if the connection just implements the + // net.PacketConn interface and the SetWriteBuffer method. + // We have no way of checking if increasing the buffer size actually worked. + if syscallConn == nil { + return conn.SetWriteBuffer(protocol.DesiredSendBufferSize) + } + + size, err := inspectWriteBuffer(syscallConn) + if err != nil { + return fmt.Errorf("failed to determine send buffer size: %w", err) + } + if size >= protocol.DesiredSendBufferSize { + utils.DefaultLogger.Debugf("Conn has send buffer of %d kiB (wanted: at least %d kiB)", size/1024, protocol.DesiredSendBufferSize/1024) + return nil + } + // Ignore the error. We check if we succeeded by querying the buffer size afterward. + _ = conn.SetWriteBuffer(protocol.DesiredSendBufferSize) + newSize, err := inspectWriteBuffer(syscallConn) + if newSize < protocol.DesiredSendBufferSize { + // Try again with RCVBUFFORCE on Linux + _ = forceSetSendBuffer(syscallConn, protocol.DesiredSendBufferSize) + newSize, err = inspectWriteBuffer(syscallConn) + if err != nil { + return fmt.Errorf("failed to determine send buffer size: %w", err) + } + } + if err != nil { + return fmt.Errorf("failed to determine send buffer size: %w", err) + } + if newSize == size { + return fmt.Errorf("failed to increase send buffer size (wanted: %d kiB, got %d kiB)", protocol.DesiredSendBufferSize/1024, newSize/1024) + } + if newSize < protocol.DesiredSendBufferSize { + return fmt.Errorf("failed to sufficiently increase send buffer size (was: %d kiB, wanted: %d kiB, got: %d kiB)", size/1024, protocol.DesiredSendBufferSize/1024, newSize/1024) + } + utils.DefaultLogger.Debugf("Increased send buffer size to %d kiB", newSize/1024) + return nil +} diff --git a/third_party/quic-go/sys_conn_df.go b/third_party/quic-go/sys_conn_df.go new file mode 100644 index 0000000..0db6150 --- /dev/null +++ b/third_party/quic-go/sys_conn_df.go @@ -0,0 +1,22 @@ +//go:build !linux && !windows && !darwin + +package quic + +import ( + "syscall" +) + +func setDF(syscall.RawConn) (bool, error) { + // no-op on unsupported platforms + return false, nil +} + +func isSendMsgSizeErr(err error) bool { + // to be implemented for more specific platforms + return false +} + +func isRecvMsgSizeErr(err error) bool { + // to be implemented for more specific platforms + return false +} diff --git a/third_party/quic-go/sys_conn_df_darwin.go b/third_party/quic-go/sys_conn_df_darwin.go new file mode 100644 index 0000000..afaa684 --- /dev/null +++ b/third_party/quic-go/sys_conn_df_darwin.go @@ -0,0 +1,90 @@ +//go:build darwin + +package quic + +import ( + "errors" + "fmt" + "strconv" + "strings" + "syscall" + + "golang.org/x/sys/unix" +) + +// for macOS versions, see https://en.wikipedia.org/wiki/Darwin_(operating_system)#Darwin_20_onwards +const ( + macOSVersion11 = 20 + macOSVersion15 = 24 +) + +func setDF(rawConn syscall.RawConn) (bool, error) { + // Setting DF bit is only supported from macOS 11. + // https://github.com/chromium/chromium/blob/117.0.5881.2/net/socket/udp_socket_posix.cc#L555 + version, err := getMacOSVersion() + if err != nil || version < macOSVersion11 { + return false, err + } + + var controlErr error + var disableDF bool + if err := rawConn.Control(func(fd uintptr) { + addr, err := unix.Getsockname(int(fd)) + if err != nil { + controlErr = fmt.Errorf("getsockname: %w", err) + return + } + + // Dual-stack sockets are effectively IPv6 sockets (with IPV6_ONLY set to 0). + // On macOS, the DF bit on dual-stack sockets is controlled by the IPV6_DONTFRAG option. + // See https://datatracker.ietf.org/doc/draft-seemann-tsvwg-udp-fragmentation/ for details. + switch addr.(type) { + case *unix.SockaddrInet4: + controlErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_DONTFRAG, 1) + case *unix.SockaddrInet6: + controlErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_DONTFRAG, 1) + + // Setting the DF bit on dual-stack sockets works since macOS Sequoia. + // Disable DF on dual-stack sockets before Sequoia. + if version < macOSVersion15 { + // check if this is a dual-stack socket by reading the IPV6_V6ONLY flag + v6only, err := unix.GetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_V6ONLY) + if err != nil { + controlErr = fmt.Errorf("getting IPV6_V6ONLY: %w", err) + return + } + disableDF = v6only == 0 + } + default: + controlErr = fmt.Errorf("unknown address type: %T", addr) + } + }); err != nil { + return false, err + } + if controlErr != nil { + return false, controlErr + } + return !disableDF, nil +} + +func isSendMsgSizeErr(err error) bool { + return errors.Is(err, unix.EMSGSIZE) +} + +func isRecvMsgSizeErr(error) bool { return false } + +func getMacOSVersion() (int, error) { + uname := &unix.Utsname{} + if err := unix.Uname(uname); err != nil { + return 0, err + } + before, _, ok := strings.Cut(string(uname.Release[:]), ".") + if !ok { + return 0, nil + } + version, err := strconv.Atoi(before) + if err != nil { + return 0, err + } + return version, nil +} diff --git a/third_party/quic-go/sys_conn_df_darwin_test.go b/third_party/quic-go/sys_conn_df_darwin_test.go new file mode 100644 index 0000000..57c5a05 --- /dev/null +++ b/third_party/quic-go/sys_conn_df_darwin_test.go @@ -0,0 +1,101 @@ +//go:build darwin + +package quic + +import ( + "net" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestIPFragmentation(t *testing.T) { + sink, err := net.ListenUDP("udp", &net.UDPAddr{Port: 0}) + require.NoError(t, err) + t.Cleanup(func() { sink.Close() }) + sinkPort := sink.LocalAddr().(*net.UDPAddr).Port + + canSendIPv4 := func(conn *net.UDPConn) bool { + _, err := conn.WriteTo([]byte("hello"), &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: sinkPort}) + return err == nil + } + + canSendIPv6 := func(conn *net.UDPConn) bool { + _, err := conn.WriteTo([]byte("hello"), &net.UDPAddr{IP: net.IPv6loopback, Port: sinkPort}) + return err == nil + } + + t.Run("udp4", func(t *testing.T) { + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + defer conn.Close() + + require.True(t, canSendIPv4(conn)) + require.False(t, canSendIPv6(conn)) + + raw, err := conn.SyscallConn() + require.NoError(t, err) + canDF, _ := setDF(raw) + require.True(t, canDF) + }) + + t.Run("udp6", func(t *testing.T) { + conn, err := net.ListenUDP("udp6", &net.UDPAddr{IP: net.IPv6loopback, Port: 0}) + require.NoError(t, err) + defer conn.Close() + + require.False(t, canSendIPv4(conn)) + require.True(t, canSendIPv6(conn)) + + raw, err := conn.SyscallConn() + require.NoError(t, err) + canDF, _ := setDF(raw) + require.True(t, canDF) + }) + + t.Run("udp, dual-stack", func(t *testing.T) { + if version, err := getMacOSVersion(); err != nil || version < macOSVersion15 { + t.Skipf("skipping on darwin %d", version-9) + } + + conn, err := net.ListenUDP("udp", &net.UDPAddr{Port: 0}) + require.NoError(t, err) + defer conn.Close() + + require.True(t, canSendIPv4(conn)) + require.True(t, canSendIPv6(conn)) + + raw, err := conn.SyscallConn() + require.NoError(t, err) + canDF, _ := setDF(raw) + require.True(t, canDF) + }) + + t.Run("udp, listening on IPv4", func(t *testing.T) { + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + defer conn.Close() + + require.True(t, canSendIPv4(conn)) + require.False(t, canSendIPv6(conn)) + + raw, err := conn.SyscallConn() + require.NoError(t, err) + canDF, _ := setDF(raw) + require.True(t, canDF) + }) + + t.Run("udp, listening on IPv6", func(t *testing.T) { + conn, err := net.ListenUDP("udp6", &net.UDPAddr{IP: net.IPv6loopback, Port: 0}) + require.NoError(t, err) + defer conn.Close() + + require.False(t, canSendIPv4(conn)) + require.True(t, canSendIPv6(conn)) + + raw, err := conn.SyscallConn() + require.NoError(t, err) + canDF, _ := setDF(raw) + require.True(t, canDF) + }) +} diff --git a/third_party/quic-go/sys_conn_df_linux.go b/third_party/quic-go/sys_conn_df_linux.go new file mode 100644 index 0000000..3bae82c --- /dev/null +++ b/third_party/quic-go/sys_conn_df_linux.go @@ -0,0 +1,42 @@ +//go:build linux + +package quic + +import ( + "errors" + "syscall" + + "golang.org/x/sys/unix" + + "github.com/apernet/quic-go/internal/utils" +) + +func setDF(rawConn syscall.RawConn) (bool, error) { + // Enabling IP_MTU_DISCOVER will force the kernel to return "sendto: message too long" + // and the datagram will not be fragmented + var errDFIPv4, errDFIPv6 error + if err := rawConn.Control(func(fd uintptr) { + errDFIPv4 = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_MTU_DISCOVER, unix.IP_PMTUDISC_PROBE) + errDFIPv6 = unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_MTU_DISCOVER, unix.IPV6_PMTUDISC_PROBE) + }); err != nil { + return false, err + } + switch { + case errDFIPv4 == nil && errDFIPv6 == nil: + utils.DefaultLogger.Debugf("Setting DF for IPv4 and IPv6.") + case errDFIPv4 == nil && errDFIPv6 != nil: + utils.DefaultLogger.Debugf("Setting DF for IPv4.") + case errDFIPv4 != nil && errDFIPv6 == nil: + utils.DefaultLogger.Debugf("Setting DF for IPv6.") + case errDFIPv4 != nil && errDFIPv6 != nil: + utils.DefaultLogger.Debugf("Setting DF failed for both IPv4 and IPv6.") + } + return true, nil +} + +func isSendMsgSizeErr(err error) bool { + // https://man7.org/linux/man-pages/man7/udp.7.html + return errors.Is(err, unix.EMSGSIZE) +} + +func isRecvMsgSizeErr(error) bool { return false } diff --git a/third_party/quic-go/sys_conn_df_windows.go b/third_party/quic-go/sys_conn_df_windows.go new file mode 100644 index 0000000..2ec00a6 --- /dev/null +++ b/third_party/quic-go/sys_conn_df_windows.go @@ -0,0 +1,52 @@ +//go:build windows + +package quic + +import ( + "errors" + "syscall" + + "golang.org/x/sys/windows" + + "github.com/apernet/quic-go/internal/utils" +) + +const ( + // https://microsoft.github.io/windows-docs-rs/doc/windows/Win32/Networking/WinSock/constant.IP_DONTFRAGMENT.html + //nolint:stylecheck + IP_DONTFRAGMENT = 14 + // https://microsoft.github.io/windows-docs-rs/doc/windows/Win32/Networking/WinSock/constant.IPV6_DONTFRAG.html + //nolint:stylecheck + IPV6_DONTFRAG = 14 +) + +func setDF(rawConn syscall.RawConn) (bool, error) { + var errDFIPv4, errDFIPv6 error + if err := rawConn.Control(func(fd uintptr) { + errDFIPv4 = windows.SetsockoptInt(windows.Handle(fd), windows.IPPROTO_IP, IP_DONTFRAGMENT, 1) + errDFIPv6 = windows.SetsockoptInt(windows.Handle(fd), windows.IPPROTO_IPV6, IPV6_DONTFRAG, 1) + }); err != nil { + return false, err + } + switch { + case errDFIPv4 == nil && errDFIPv6 == nil: + utils.DefaultLogger.Debugf("Setting DF for IPv4 and IPv6.") + case errDFIPv4 == nil && errDFIPv6 != nil: + utils.DefaultLogger.Debugf("Setting DF for IPv4.") + case errDFIPv4 != nil && errDFIPv6 == nil: + utils.DefaultLogger.Debugf("Setting DF for IPv6.") + case errDFIPv4 != nil && errDFIPv6 != nil: + utils.DefaultLogger.Debugf("Setting DF failed for both IPv4 and IPv6.") + } + return true, nil +} + +func isSendMsgSizeErr(err error) bool { + // https://docs.microsoft.com/en-us/windows/win32/winsock/windows-sockets-error-codes-2 + return errors.Is(err, windows.WSAEMSGSIZE) +} + +func isRecvMsgSizeErr(err error) bool { + // https://docs.microsoft.com/en-us/windows/win32/winsock/windows-sockets-error-codes-2 + return errors.Is(err, windows.WSAEMSGSIZE) +} diff --git a/third_party/quic-go/sys_conn_helper_darwin.go b/third_party/quic-go/sys_conn_helper_darwin.go new file mode 100644 index 0000000..a04bfb3 --- /dev/null +++ b/third_party/quic-go/sys_conn_helper_darwin.go @@ -0,0 +1,38 @@ +//go:build darwin + +package quic + +import ( + "encoding/binary" + "net/netip" + "syscall" + + "golang.org/x/sys/unix" +) + +const ( + msgTypeIPTOS = unix.IP_RECVTOS + ipv4PKTINFO = unix.IP_RECVPKTINFO +) + +const ecnIPv4DataLen = 4 + +// ReadBatch only returns a single packet on OSX, +// see https://godoc.org/golang.org/x/net/ipv4#PacketConn.ReadBatch. +const batchSize = 1 + +func parseIPv4PktInfo(body []byte) (ip netip.Addr, ifIndex uint32, ok bool) { + // struct in_pktinfo { + // unsigned int ipi_ifindex; /* Interface index */ + // struct in_addr ipi_spec_dst; /* Local address */ + // struct in_addr ipi_addr; /* Header Destination address */ + // }; + if len(body) != 12 { + return netip.Addr{}, 0, false + } + return netip.AddrFrom4(*(*[4]byte)(body[8:12])), binary.NativeEndian.Uint32(body), true +} + +func isGSOEnabled(syscall.RawConn) bool { return false } + +func isECNEnabled() bool { return !isECNDisabledUsingEnv() } diff --git a/third_party/quic-go/sys_conn_helper_freebsd.go b/third_party/quic-go/sys_conn_helper_freebsd.go new file mode 100644 index 0000000..521f80d --- /dev/null +++ b/third_party/quic-go/sys_conn_helper_freebsd.go @@ -0,0 +1,33 @@ +//go:build freebsd + +package quic + +import ( + "net/netip" + "syscall" + + "golang.org/x/sys/unix" +) + +const ( + msgTypeIPTOS = unix.IP_RECVTOS + ipv4PKTINFO = 0x7 +) + +const ecnIPv4DataLen = 1 + +const batchSize = 8 + +func parseIPv4PktInfo(body []byte) (ip netip.Addr, _ uint32, ok bool) { + // struct in_pktinfo { + // struct in_addr ipi_addr; /* Header Destination address */ + // }; + if len(body) != 4 { + return netip.Addr{}, 0, false + } + return netip.AddrFrom4(*(*[4]byte)(body)), 0, true +} + +func isGSOEnabled(syscall.RawConn) bool { return false } + +func isECNEnabled() bool { return !isECNDisabledUsingEnv() } diff --git a/third_party/quic-go/sys_conn_helper_linux.go b/third_party/quic-go/sys_conn_helper_linux.go new file mode 100644 index 0000000..4017754 --- /dev/null +++ b/third_party/quic-go/sys_conn_helper_linux.go @@ -0,0 +1,156 @@ +//go:build linux + +package quic + +import ( + "encoding/binary" + "errors" + "net/netip" + "os" + "strconv" + "syscall" + "unsafe" + + "golang.org/x/sys/unix" +) + +const ( + msgTypeIPTOS = unix.IP_TOS + ipv4PKTINFO = unix.IP_PKTINFO +) + +const ecnIPv4DataLen = 1 + +const batchSize = 8 // needs to smaller than MaxUint8 (otherwise the type of oobConn.readPos has to be changed) + +var kernelVersionMajor int + +func init() { + kernelVersionMajor, _ = kernelVersion() +} + +func forceSetReceiveBuffer(c syscall.RawConn, bytes int) error { + var serr error + if err := c.Control(func(fd uintptr) { + serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, bytes) + }); err != nil { + return err + } + return serr +} + +func forceSetSendBuffer(c syscall.RawConn, bytes int) error { + var serr error + if err := c.Control(func(fd uintptr) { + serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, bytes) + }); err != nil { + return err + } + return serr +} + +func parseIPv4PktInfo(body []byte) (ip netip.Addr, ifIndex uint32, ok bool) { + // struct in_pktinfo { + // unsigned int ipi_ifindex; /* Interface index */ + // struct in_addr ipi_spec_dst; /* Local address */ + // struct in_addr ipi_addr; /* Header Destination address */ + // }; + if len(body) != 12 { + return netip.Addr{}, 0, false + } + return netip.AddrFrom4(*(*[4]byte)(body[8:12])), binary.NativeEndian.Uint32(body), true +} + +// isGSOEnabled tests if the kernel supports GSO. +// Sending with GSO might still fail later on, if the interface doesn't support it (see isGSOError). +func isGSOEnabled(conn syscall.RawConn) bool { + if kernelVersionMajor < 5 { + return false + } + disabled, err := strconv.ParseBool(os.Getenv("QUIC_GO_DISABLE_GSO")) + if err == nil && disabled { + return false + } + var serr error + if err := conn.Control(func(fd uintptr) { + _, serr = unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_SEGMENT) + }); err != nil { + return false + } + return serr == nil +} + +func appendUDPSegmentSizeMsg(b []byte, size uint16) []byte { + startLen := len(b) + const dataLen = 2 // payload is a uint16 + b = append(b, make([]byte, unix.CmsgSpace(dataLen))...) + h := (*unix.Cmsghdr)(unsafe.Pointer(&b[startLen])) + h.Level = syscall.IPPROTO_UDP + h.Type = unix.UDP_SEGMENT + h.SetLen(unix.CmsgLen(dataLen)) + + // UnixRights uses the private `data` method, but I *think* this achieves the same goal. + offset := startLen + unix.CmsgSpace(0) + *(*uint16)(unsafe.Pointer(&b[offset])) = size + return b +} + +func isGSOError(err error) bool { + var serr *os.SyscallError + if errors.As(err, &serr) { + // EIO is returned by udp_send_skb() if the device driver does not have tx checksums enabled, + // which is a hard requirement of UDP_SEGMENT. See: + // https://git.kernel.org/pub/scm/docs/man-pages/man-pages.git/tree/man7/udp.7?id=806eabd74910447f21005160e90957bde4db0183#n228 + // https://git.kernel.org/pub/scm/linux/kernel/git/torvalds/linux.git/tree/net/ipv4/udp.c?h=v6.2&id=c9c3395d5e3dcc6daee66c6908354d47bf98cb0c#n942 + return serr.Err == unix.EIO || serr.Err == unix.EINVAL + } + return false +} + +// The first sendmsg call on a new UDP socket sometimes errors on Linux. +// It's not clear why this happens. +// See https://github.com/golang/go/issues/63322. +func isPermissionError(err error) bool { + var serr *os.SyscallError + if errors.As(err, &serr) { + return serr.Syscall == "sendmsg" && serr.Err == unix.EPERM + } + return false +} + +func isECNEnabled() bool { + return kernelVersionMajor >= 5 && !isECNDisabledUsingEnv() +} + +// kernelVersion returns major and minor kernel version numbers, parsed from +// the syscall.Uname's Release field, or 0, 0 if the version can't be obtained +// or parsed. +// +// copied from the standard library's internal/syscall/unix/kernel_version_linux.go +func kernelVersion() (major, minor int) { + var uname syscall.Utsname + if err := syscall.Uname(&uname); err != nil { + return + } + + var ( + values [2]int + value, vi int + ) + for _, c := range uname.Release { + if '0' <= c && c <= '9' { + value = (value * 10) + int(c-'0') + } else { + // Note that we're assuming N.N.N here. + // If we see anything else, we are likely to mis-parse it. + values[vi] = value + vi++ + if vi >= len(values) { + break + } + value = 0 + } + } + + return values[0], values[1] +} diff --git a/third_party/quic-go/sys_conn_helper_linux_test.go b/third_party/quic-go/sys_conn_helper_linux_test.go new file mode 100644 index 0000000..547d320 --- /dev/null +++ b/third_party/quic-go/sys_conn_helper_linux_test.go @@ -0,0 +1,79 @@ +//go:build linux + +package quic + +import ( + "errors" + "net" + "os" + "testing" + + "golang.org/x/sys/unix" + + "github.com/stretchr/testify/require" +) + +var ( + errGSO = &os.SyscallError{Err: unix.EIO} + errNotPermitted = &os.SyscallError{Syscall: "sendmsg", Err: unix.EPERM} +) + +func TestForcingReceiveBufferSize(t *testing.T) { + if os.Getuid() != 0 { + t.Skip("Must be root to force change the receive buffer size") + } + + c, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + defer c.Close() + syscallConn, err := c.(*net.UDPConn).SyscallConn() + require.NoError(t, err) + + const small = 256 << 10 // 256 KB + require.NoError(t, forceSetReceiveBuffer(syscallConn, small)) + + size, err := inspectReadBuffer(syscallConn) + require.NoError(t, err) + // the kernel doubles this value (to allow space for bookkeeping overhead) + require.Equal(t, 2*small, size) + + const large = 32 << 20 // 32 MB + require.NoError(t, forceSetReceiveBuffer(syscallConn, large)) + size, err = inspectReadBuffer(syscallConn) + require.NoError(t, err) + // the kernel doubles this value (to allow space for bookkeeping overhead) + require.Equal(t, 2*large, size) +} + +func TestForcingSendBufferSize(t *testing.T) { + if os.Getuid() != 0 { + t.Skip("Must be root to force change the send buffer size") + } + + c, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + defer c.Close() + syscallConn, err := c.(*net.UDPConn).SyscallConn() + require.NoError(t, err) + + const small = 256 << 10 // 256 KB + require.NoError(t, forceSetSendBuffer(syscallConn, small)) + + size, err := inspectWriteBuffer(syscallConn) + require.NoError(t, err) + // the kernel doubles this value (to allow space for bookkeeping overhead) + require.Equal(t, 2*small, size) + + const large = 32 << 20 // 32 MB + require.NoError(t, forceSetSendBuffer(syscallConn, large)) + size, err = inspectWriteBuffer(syscallConn) + require.NoError(t, err) + // the kernel doubles this value (to allow space for bookkeeping overhead) + require.Equal(t, 2*large, size) +} + +func TestGSOError(t *testing.T) { + require.True(t, isGSOError(errGSO)) + require.False(t, isGSOError(nil)) + require.False(t, isGSOError(errors.New("test"))) +} diff --git a/third_party/quic-go/sys_conn_helper_nonlinux.go b/third_party/quic-go/sys_conn_helper_nonlinux.go new file mode 100644 index 0000000..f8d6980 --- /dev/null +++ b/third_party/quic-go/sys_conn_helper_nonlinux.go @@ -0,0 +1,10 @@ +//go:build !linux + +package quic + +func forceSetReceiveBuffer(c any, bytes int) error { return nil } +func forceSetSendBuffer(c any, bytes int) error { return nil } + +func appendUDPSegmentSizeMsg([]byte, uint16) []byte { return nil } +func isGSOError(error) bool { return false } +func isPermissionError(err error) bool { return false } diff --git a/third_party/quic-go/sys_conn_helper_nonlinux_test.go b/third_party/quic-go/sys_conn_helper_nonlinux_test.go new file mode 100644 index 0000000..0967124 --- /dev/null +++ b/third_party/quic-go/sys_conn_helper_nonlinux_test.go @@ -0,0 +1,10 @@ +//go:build !linux + +package quic + +import "errors" + +var ( + errGSO = errors.New("fake GSO error") + errNotPermitted = errors.New("fake not permitted error") +) diff --git a/third_party/quic-go/sys_conn_no_oob.go b/third_party/quic-go/sys_conn_no_oob.go new file mode 100644 index 0000000..57916c8 --- /dev/null +++ b/third_party/quic-go/sys_conn_no_oob.go @@ -0,0 +1,21 @@ +//go:build !darwin && !linux && !freebsd && !windows + +package quic + +import ( + "net" + "net/netip" +) + +func newConn(c net.PacketConn, supportsDF bool, _ bool) (*basicConn, error) { + return &basicConn{PacketConn: c, supportsDF: supportsDF}, nil +} + +func inspectReadBuffer(any) (int, error) { return 0, nil } +func inspectWriteBuffer(any) (int, error) { return 0, nil } + +type packetInfo struct { + addr netip.Addr +} + +func (i *packetInfo) OOB() []byte { return nil } diff --git a/third_party/quic-go/sys_conn_oob.go b/third_party/quic-go/sys_conn_oob.go new file mode 100644 index 0000000..1a56d3c --- /dev/null +++ b/third_party/quic-go/sys_conn_oob.go @@ -0,0 +1,344 @@ +//go:build darwin || linux || freebsd + +package quic + +import ( + "encoding/binary" + "errors" + "log" + "net" + "net/netip" + "os" + "strconv" + "sync" + "syscall" + "unsafe" + + "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" + "golang.org/x/sys/unix" + + "github.com/apernet/quic-go/internal/monotime" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" +) + +const ( + ecnMask = 0x3 + oobBufferSize = 128 +) + +// Contrary to what the naming suggests, the ipv{4,6}.Message is not dependent on the IP version. +// They're both just aliases for x/net/internal/socket.Message. +// This means we can use this struct to read from a socket that receives both IPv4 and IPv6 messages. +var _ ipv4.Message = ipv6.Message{} + +type batchConn interface { + ReadBatch(ms []ipv4.Message, flags int) (int, error) +} + +func inspectReadBuffer(c syscall.RawConn) (int, error) { + var size int + var serr error + if err := c.Control(func(fd uintptr) { + size, serr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_RCVBUF) + }); err != nil { + return 0, err + } + return size, serr +} + +func inspectWriteBuffer(c syscall.RawConn) (int, error) { + var size int + var serr error + if err := c.Control(func(fd uintptr) { + size, serr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_SNDBUF) + }); err != nil { + return 0, err + } + return size, serr +} + +func isECNDisabledUsingEnv() bool { + disabled, err := strconv.ParseBool(os.Getenv("QUIC_GO_DISABLE_ECN")) + return err == nil && disabled +} + +type oobConn struct { + OOBCapablePacketConn + batchConn batchConn + + readPos uint8 + // Packets received from the kernel, but not yet returned by ReadPacket(). + messages []ipv4.Message + buffers [batchSize]*packetBuffer + + cap connCapabilities +} + +var _ rawConn = &oobConn{} + +func newConn(c OOBCapablePacketConn, supportsDF bool, disableGSO bool) (*oobConn, error) { + rawConn, err := c.SyscallConn() + if err != nil { + return nil, err + } + var needsPacketInfo bool + if udpAddr, ok := c.LocalAddr().(*net.UDPAddr); ok && udpAddr.IP.IsUnspecified() { + needsPacketInfo = true + } + // We don't know if this a IPv4-only, IPv6-only or a IPv4-and-IPv6 connection. + // Try enabling receiving of ECN and packet info for both IP versions. + // We expect at least one of those syscalls to succeed. + var errECNIPv4, errECNIPv6, errPIIPv4, errPIIPv6 error + if err := rawConn.Control(func(fd uintptr) { + errECNIPv4 = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1) + errECNIPv6 = unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_RECVTCLASS, 1) + + if needsPacketInfo { + errPIIPv4 = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, ipv4PKTINFO, 1) + errPIIPv6 = unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_RECVPKTINFO, 1) + } + }); err != nil { + return nil, err + } + switch { + case errECNIPv4 == nil && errECNIPv6 == nil: + utils.DefaultLogger.Debugf("Activating reading of ECN bits for IPv4 and IPv6.") + case errECNIPv4 == nil && errECNIPv6 != nil: + utils.DefaultLogger.Debugf("Activating reading of ECN bits for IPv4.") + case errECNIPv4 != nil && errECNIPv6 == nil: + utils.DefaultLogger.Debugf("Activating reading of ECN bits for IPv6.") + case errECNIPv4 != nil && errECNIPv6 != nil: + return nil, errors.New("activating ECN failed for both IPv4 and IPv6") + } + if needsPacketInfo { + switch { + case errPIIPv4 == nil && errPIIPv6 == nil: + utils.DefaultLogger.Debugf("Activating reading of packet info for IPv4 and IPv6.") + case errPIIPv4 == nil && errPIIPv6 != nil: + utils.DefaultLogger.Debugf("Activating reading of packet info bits for IPv4.") + case errPIIPv4 != nil && errPIIPv6 == nil: + utils.DefaultLogger.Debugf("Activating reading of packet info bits for IPv6.") + case errPIIPv4 != nil && errPIIPv6 != nil: + return nil, errors.New("activating packet info failed for both IPv4 and IPv6") + } + } + + // Allows callers to pass in a connection that already satisfies batchConn interface + // to make use of the optimisation. Otherwise, ipv4.NewPacketConn would unwrap the file descriptor + // via SyscallConn(), and read it that way, which might not be what the caller wants. + var bc batchConn + if ibc, ok := c.(batchConn); ok { + bc = ibc + } else { + bc = ipv4.NewPacketConn(c) + } + + msgs := make([]ipv4.Message, batchSize) + for i := range msgs { + // preallocate the [][]byte + msgs[i].Buffers = make([][]byte, 1) + } + oobConn := &oobConn{ + OOBCapablePacketConn: c, + batchConn: bc, + messages: msgs, + readPos: batchSize, + cap: connCapabilities{ + DF: supportsDF, + GSO: !disableGSO && isGSOEnabled(rawConn), + ECN: isECNEnabled(), + }, + } + for i := range batchSize { + oobConn.messages[i].OOB = make([]byte, oobBufferSize) + } + return oobConn, nil +} + +var invalidCmsgOnceV4, invalidCmsgOnceV6 sync.Once + +func (c *oobConn) ReadPacket() (receivedPacket, error) { + if len(c.messages) == int(c.readPos) { // all messages read. Read the next batch of messages. + c.messages = c.messages[:batchSize] + // replace buffers data buffers up to the packet that has been consumed during the last ReadBatch call + for i := uint8(0); i < c.readPos; i++ { + buffer := getPacketBuffer() + buffer.Data = buffer.Data[:protocol.MaxPacketBufferSize] + c.buffers[i] = buffer + c.messages[i].Buffers[0] = c.buffers[i].Data + } + c.readPos = 0 + + n, err := c.batchConn.ReadBatch(c.messages, 0) + if n == 0 || err != nil { + return receivedPacket{}, err + } + c.messages = c.messages[:n] + } + + msg := c.messages[c.readPos] + buffer := c.buffers[c.readPos] + c.readPos++ + + data := msg.OOB[:msg.NN] + p := receivedPacket{ + remoteAddr: msg.Addr, + rcvTime: monotime.Now(), + data: msg.Buffers[0][:msg.N], + buffer: buffer, + } + for len(data) > 0 { + hdr, body, remainder, err := unix.ParseOneSocketControlMessage(data) + if err != nil { + return receivedPacket{}, err + } + if hdr.Level == unix.IPPROTO_IP { + switch hdr.Type { + case msgTypeIPTOS: + if len(body) != 1 { + return receivedPacket{}, errors.New("invalid IPTOS size") + } + p.ecn = protocol.ParseECNHeaderBits(body[0] & ecnMask) + case ipv4PKTINFO: + ip, ifIndex, ok := parseIPv4PktInfo(body) + if ok { + p.info.addr = ip + p.info.ifIndex = ifIndex + } else { + invalidCmsgOnceV4.Do(func() { + log.Printf("Received invalid IPv4 packet info control message: %+x. "+ + "This should never occur, please open a new issue and include details about the architecture.", body) + }) + } + } + } + if hdr.Level == unix.IPPROTO_IPV6 { + switch hdr.Type { + case unix.IPV6_TCLASS: + if len(body) != 4 { + return receivedPacket{}, errors.New("invalid IPV6_TCLASS size") + } + bits := uint8(binary.NativeEndian.Uint32(body)) & ecnMask + p.ecn = protocol.ParseECNHeaderBits(bits) + case unix.IPV6_PKTINFO: + // struct in6_pktinfo { + // struct in6_addr ipi6_addr; /* src/dst IPv6 address */ + // unsigned int ipi6_ifindex; /* send/recv interface index */ + // }; + if len(body) == 20 { + p.info.addr = netip.AddrFrom16(*(*[16]byte)(body[:16])).Unmap() + p.info.ifIndex = binary.NativeEndian.Uint32(body[16:]) + } else { + invalidCmsgOnceV6.Do(func() { + log.Printf("Received invalid IPv6 packet info control message: %+x. "+ + "This should never occur, please open a new issue and include details about the architecture.", body) + }) + } + } + } + data = remainder + } + return p, nil +} + +// WritePacket writes a new packet. +func (c *oobConn) WritePacket(b []byte, addr net.Addr, packetInfoOOB []byte, gsoSize uint16, ecn protocol.ECN) (int, error) { + oob := packetInfoOOB + if gsoSize > 0 { + if !c.capabilities().GSO { + panic("GSO disabled") + } + // Only request UDP GSO when the payload will actually be segmented. + // Some drivers/devices misbehave when UDP_SEGMENT is set for an effectively + // single-segment send (segment_size >= payload length). This mirrors quinn-udp's + // behavior. + if len(b) > int(gsoSize) { + oob = appendUDPSegmentSizeMsg(oob, gsoSize) + } + } + if ecn != protocol.ECNUnsupported { + if !c.capabilities().ECN { + panic("tried to send an ECN-marked packet although ECN is disabled") + } + if remoteUDPAddr, ok := addr.(*net.UDPAddr); ok { + if remoteUDPAddr.IP.To4() != nil { + oob = appendIPv4ECNMsg(oob, ecn) + } else { + oob = appendIPv6ECNMsg(oob, ecn) + } + } + } + n, _, err := c.WriteMsgUDP(b, oob, addr.(*net.UDPAddr)) + return n, err +} + +func (c *oobConn) capabilities() connCapabilities { + return c.cap +} + +type packetInfo struct { + addr netip.Addr + ifIndex uint32 +} + +func (info *packetInfo) OOB() []byte { + if info == nil { + return nil + } + if info.addr.Is4() { + ip := info.addr.As4() + // struct in_pktinfo { + // unsigned int ipi_ifindex; /* Interface index */ + // struct in_addr ipi_spec_dst; /* Local address */ + // struct in_addr ipi_addr; /* Header Destination address */ + // }; + cm := ipv4.ControlMessage{ + Src: ip[:], + IfIndex: int(info.ifIndex), + } + return cm.Marshal() + } else if info.addr.Is6() { + ip := info.addr.As16() + // struct in6_pktinfo { + // struct in6_addr ipi6_addr; /* src/dst IPv6 address */ + // unsigned int ipi6_ifindex; /* send/recv interface index */ + // }; + cm := ipv6.ControlMessage{ + Src: ip[:], + IfIndex: int(info.ifIndex), + } + return cm.Marshal() + } + return nil +} + +func appendIPv4ECNMsg(b []byte, val protocol.ECN) []byte { + startLen := len(b) + b = append(b, make([]byte, unix.CmsgSpace(ecnIPv4DataLen))...) + h := (*unix.Cmsghdr)(unsafe.Pointer(&b[startLen])) + h.Level = syscall.IPPROTO_IP + h.Type = unix.IP_TOS + h.SetLen(unix.CmsgLen(ecnIPv4DataLen)) + + // UnixRights uses the private `data` method, but I *think* this achieves the same goal. + offset := startLen + unix.CmsgSpace(0) + b[offset] = val.ToHeaderBits() + return b +} + +func appendIPv6ECNMsg(b []byte, val protocol.ECN) []byte { + startLen := len(b) + const dataLen = 4 + b = append(b, make([]byte, unix.CmsgSpace(dataLen))...) + h := (*unix.Cmsghdr)(unsafe.Pointer(&b[startLen])) + h.Level = syscall.IPPROTO_IPV6 + h.Type = unix.IPV6_TCLASS + h.SetLen(unix.CmsgLen(dataLen)) + + // UnixRights uses the private `data` method, but I *think* this achieves the same goal. + offset := startLen + unix.CmsgSpace(0) + binary.NativeEndian.PutUint32(b[offset:offset+dataLen], uint32(val.ToHeaderBits())) + return b +} diff --git a/third_party/quic-go/sys_conn_oob_test.go b/third_party/quic-go/sys_conn_oob_test.go new file mode 100644 index 0000000..34d6022 --- /dev/null +++ b/third_party/quic-go/sys_conn_oob_test.go @@ -0,0 +1,334 @@ +//go:build darwin || linux || freebsd + +package quic + +import ( + "fmt" + "net" + "testing" + "time" + + "golang.org/x/net/ipv4" + "golang.org/x/sys/unix" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" +) + +func isIPv4(ip net.IP) bool { return ip.To4() != nil } + +func runSysConnServer(t *testing.T, network string, addr *net.UDPAddr) (*net.UDPAddr, <-chan receivedPacket) { + t.Helper() + udpConn, err := net.ListenUDP(network, addr) + require.NoError(t, err) + t.Cleanup(func() { udpConn.Close() }) + + oobConn, err := newConn(udpConn, true, false) + require.NoError(t, err) + require.True(t, oobConn.capabilities().DF) + + packetChan := make(chan receivedPacket, 1) + go func() { + for { + p, err := oobConn.ReadPacket() + if err != nil { + return + } + packetChan <- p + } + }() + return udpConn.LocalAddr().(*net.UDPAddr), packetChan +} + +// sendUDPPacketWithECN opens a new UDP socket and sends one packet with the ECN set. +// It returns the local address of the socket. +func sendUDPPacketWithECN(t *testing.T, network string, addr *net.UDPAddr, setECN func(uintptr)) net.Addr { + conn, err := net.DialUDP(network, nil, addr) + require.NoError(t, err) + t.Cleanup(func() { conn.Close() }) + + rawConn, err := conn.SyscallConn() + require.NoError(t, err) + require.NoError(t, rawConn.Control(func(fd uintptr) { setECN(fd) })) + _, err = conn.Write([]byte("foobar")) + require.NoError(t, err) + return conn.LocalAddr() +} + +func TestReadECNFlagsIPv4(t *testing.T) { + addr, packetChan := runSysConnServer(t, "udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + + sentFrom := sendUDPPacketWithECN(t, + "udp4", + addr, + func(fd uintptr) { + require.NoError(t, unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_TOS, 2)) + }, + ) + + select { + case p := <-packetChan: + require.WithinDuration(t, time.Now(), p.rcvTime.ToTime(), scaleDuration(20*time.Millisecond)) + require.Equal(t, []byte("foobar"), p.data) + require.Equal(t, sentFrom, p.remoteAddr) + require.Equal(t, protocol.ECT0, p.ecn) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } +} + +func TestReadECNFlagsIPv6(t *testing.T) { + addr, packetChan := runSysConnServer(t, "udp6", &net.UDPAddr{IP: net.IPv6loopback, Port: 0}) + + sentFrom := sendUDPPacketWithECN(t, + "udp6", + addr, + func(fd uintptr) { + require.NoError(t, unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_TCLASS, 3)) + }, + ) + + select { + case p := <-packetChan: + require.WithinDuration(t, time.Now(), p.rcvTime.ToTime(), scaleDuration(20*time.Millisecond)) + require.Equal(t, []byte("foobar"), p.data) + require.Equal(t, sentFrom, p.remoteAddr) + require.Equal(t, protocol.ECNCE, p.ecn) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } +} + +func TestReadECNFlagsDualStack(t *testing.T) { + addr, packetChan := runSysConnServer(t, "udp", &net.UDPAddr{IP: net.IPv4(0, 0, 0, 0), Port: 0}) + + // IPv4 + sentFrom := sendUDPPacketWithECN(t, + "udp4", + &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: addr.Port}, + func(fd uintptr) { + require.NoError(t, unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_TOS, 3)) + }, + ) + + select { + case p := <-packetChan: + require.True(t, isIPv4(p.remoteAddr.(*net.UDPAddr).IP)) + require.Equal(t, sentFrom.String(), p.remoteAddr.String()) + require.Equal(t, protocol.ECNCE, p.ecn) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } + + // IPv6 + sentFrom = sendUDPPacketWithECN(t, + "udp6", + &net.UDPAddr{IP: net.IPv6loopback, Port: addr.Port}, + func(fd uintptr) { + require.NoError(t, unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_TCLASS, 1)) + }, + ) + + select { + case p := <-packetChan: + require.Equal(t, sentFrom, p.remoteAddr) + require.False(t, isIPv4(p.remoteAddr.(*net.UDPAddr).IP)) + require.Equal(t, protocol.ECT1, p.ecn) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } +} + +func TestSendPacketsWithECNOnIPv4(t *testing.T) { + addr, packetChan := runSysConnServer(t, "udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + + c, err := net.ListenUDP("udp4", nil) + require.NoError(t, err) + defer c.Close() + + for _, val := range []protocol.ECN{protocol.ECNNon, protocol.ECT1, protocol.ECT0, protocol.ECNCE} { + _, _, err = c.WriteMsgUDP([]byte("foobar"), appendIPv4ECNMsg([]byte{}, val), addr) + require.NoError(t, err) + select { + case p := <-packetChan: + require.Equal(t, []byte("foobar"), p.data) + require.Equal(t, val, p.ecn) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } + } +} + +func TestSendPacketsWithECNOnIPv6(t *testing.T) { + addr, packetChan := runSysConnServer(t, "udp6", &net.UDPAddr{IP: net.IPv6loopback, Port: 0}) + + c, err := net.ListenUDP("udp6", nil) + require.NoError(t, err) + defer c.Close() + + for _, val := range []protocol.ECN{protocol.ECNNon, protocol.ECT1, protocol.ECT0, protocol.ECNCE} { + _, _, err = c.WriteMsgUDP([]byte("foobar"), appendIPv6ECNMsg([]byte{}, val), addr) + require.NoError(t, err) + select { + case p := <-packetChan: + require.Equal(t, []byte("foobar"), p.data) + require.Equal(t, val, p.ecn) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } + } +} + +func TestSysConnPacketInfoIPv4(t *testing.T) { + // need to listen on 0.0.0.0, otherwise we won't get the packet info + addr, packetChan := runSysConnServer(t, "udp4", &net.UDPAddr{IP: net.IPv4zero, Port: 0}) + + conn, err := net.DialUDP("udp4", nil, addr) + require.NoError(t, err) + defer conn.Close() + _, err = conn.Write([]byte("foobar")) + require.NoError(t, err) + + select { + case p := <-packetChan: + require.WithinDuration(t, time.Now(), p.rcvTime.ToTime(), scaleDuration(50*time.Millisecond)) + require.Equal(t, []byte("foobar"), p.data) + require.Equal(t, conn.LocalAddr(), p.remoteAddr) + require.True(t, p.info.addr.IsValid()) + require.True(t, isIPv4(p.info.addr.AsSlice())) + require.Equal(t, net.IPv4(127, 0, 0, 1).String(), p.info.addr.String()) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } +} + +func TestSysConnPacketInfoIPv6(t *testing.T) { + // need to listen on ::, otherwise we won't get the packet info + addr, packetChan := runSysConnServer(t, "udp6", &net.UDPAddr{IP: net.IPv6zero, Port: 0}) + + conn, err := net.DialUDP("udp6", nil, addr) + require.NoError(t, err) + defer conn.Close() + _, err = conn.Write([]byte("foobar")) + require.NoError(t, err) + + select { + case p := <-packetChan: + require.WithinDuration(t, time.Now(), p.rcvTime.ToTime(), scaleDuration(20*time.Millisecond)) + require.Equal(t, []byte("foobar"), p.data) + require.Equal(t, conn.LocalAddr(), p.remoteAddr) + require.NotNil(t, p.info) + require.Equal(t, net.IPv6loopback, net.IP(p.info.addr.AsSlice())) + case <-time.After(time.Second): + t.Fatal("timeout waiting for packet") + } +} + +func TestSysConnPacketInfoDualStack(t *testing.T) { + addr, packetChan := runSysConnServer(t, "udp", &net.UDPAddr{}) + + // IPv4 + conn4, err := net.DialUDP("udp4", nil, &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: addr.Port}) + require.NoError(t, err) + defer conn4.Close() + _, err = conn4.Write([]byte("foobar")) + require.NoError(t, err) + + select { + case p := <-packetChan: + require.True(t, isIPv4(p.remoteAddr.(*net.UDPAddr).IP)) + require.NotNil(t, p.info) + require.True(t, p.info.addr.Is4()) + require.Equal(t, net.IPv4(127, 0, 0, 1).String(), p.info.addr.String()) + case <-time.After(time.Second): + t.Fatal("timeout waiting for IPv4 packet") + } + + // IPv6 + conn6, err := net.DialUDP("udp6", nil, addr) + require.NoError(t, err) + defer conn6.Close() + _, err = conn6.Write([]byte("foobar")) + require.NoError(t, err) + + select { + case p := <-packetChan: + require.False(t, isIPv4(p.remoteAddr.(*net.UDPAddr).IP)) + require.NotNil(t, p.info) + require.Equal(t, net.IPv6loopback.String(), p.info.addr.String()) + case <-time.After(time.Second): + t.Fatal("timeout waiting for IPv6 packet") + } +} + +type oobRecordingConn struct { + *net.UDPConn + oobs [][]byte +} + +func (c *oobRecordingConn) WriteMsgUDP(b, oob []byte, addr *net.UDPAddr) (n, oobn int, err error) { + c.oobs = append(c.oobs, oob) + return c.UDPConn.WriteMsgUDP(b, oob, addr) +} + +type mockBatchConn struct { + t *testing.T + numMsgRead int + + callCounter int +} + +var _ batchConn = &mockBatchConn{} + +func (c *mockBatchConn) ReadBatch(ms []ipv4.Message, _ int) (int, error) { + require.Len(c.t, ms, batchSize) + for i := 0; i < c.numMsgRead; i++ { + require.Len(c.t, ms[i].Buffers, 1) + require.Len(c.t, ms[i].Buffers[0], protocol.MaxPacketBufferSize) + data := fmt.Appendf(nil, "message %d", c.callCounter*c.numMsgRead+i) + ms[i].Buffers[0] = data + ms[i].N = len(data) + } + c.callCounter++ + return c.numMsgRead, nil +} + +func TestReadsMultipleMessagesInOneBatch(t *testing.T) { + bc := &mockBatchConn{t: t, numMsgRead: batchSize/2 + 1} + + udpConn := newUDPConnLocalhost(t) + oobConn, err := newConn(udpConn, true, false) + require.NoError(t, err) + oobConn.batchConn = bc + + for i := range batchSize + 1 { + p, err := oobConn.ReadPacket() + require.NoError(t, err) + require.Equal(t, fmt.Sprintf("message %d", i), string(p.data)) + } + require.Equal(t, 2, bc.callCounter) +} + +func TestSysConnSendGSO(t *testing.T) { + if !platformSupportsGSO { + t.Skip("GSO not supported on this platform") + } + + udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + c := &oobRecordingConn{UDPConn: udpConn} + oobConn, err := newConn(c, true, false) + require.NoError(t, err) + require.True(t, oobConn.capabilities().GSO) + + oob := make([]byte, 0, 123) + oobConn.WritePacket([]byte("foobar"), udpConn.LocalAddr(), oob, 3, protocol.ECNCE) + require.Len(t, c.oobs, 1) + oobMsg := c.oobs[0] + require.NotEmpty(t, oobMsg) + require.Equal(t, cap(oob), cap(oobMsg)) // check that it appended to oob + expected := appendUDPSegmentSizeMsg([]byte{}, 3) + // Check that the first control message is the OOB control message. + require.Equal(t, expected, oobMsg[:len(expected)]) +} diff --git a/third_party/quic-go/sys_conn_test.go b/third_party/quic-go/sys_conn_test.go new file mode 100644 index 0000000..585a767 --- /dev/null +++ b/third_party/quic-go/sys_conn_test.go @@ -0,0 +1,32 @@ +package quic + +import ( + "net" + "testing" + "time" + + "github.com/apernet/quic-go/internal/protocol" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +func TestBasicConn(t *testing.T) { + mockCtrl := gomock.NewController(t) + + c := NewMockPacketConn(mockCtrl) + addr := &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234} + c.EXPECT().ReadFrom(gomock.Any()).DoAndReturn(func(b []byte) (int, net.Addr, error) { + data := []byte("foobar") + require.Equal(t, protocol.MaxPacketBufferSize, len(b)) + return copy(b, data), addr, nil + }) + + conn, err := wrapConn(c, false) + require.NoError(t, err) + p, err := conn.ReadPacket() + require.NoError(t, err) + require.Equal(t, []byte("foobar"), p.data) + require.WithinDuration(t, time.Now(), p.rcvTime.ToTime(), scaleDuration(100*time.Millisecond)) + require.Equal(t, addr, p.remoteAddr) +} diff --git a/third_party/quic-go/sys_conn_windows.go b/third_party/quic-go/sys_conn_windows.go new file mode 100644 index 0000000..c945dbc --- /dev/null +++ b/third_party/quic-go/sys_conn_windows.go @@ -0,0 +1,42 @@ +//go:build windows + +package quic + +import ( + "net/netip" + "syscall" + + "golang.org/x/sys/windows" +) + +func newConn(c OOBCapablePacketConn, supportsDF bool, _ bool) (*basicConn, error) { + return &basicConn{PacketConn: c, supportsDF: supportsDF}, nil +} + +func inspectReadBuffer(c syscall.RawConn) (int, error) { + var size int + var serr error + if err := c.Control(func(fd uintptr) { + size, serr = windows.GetsockoptInt(windows.Handle(fd), windows.SOL_SOCKET, windows.SO_RCVBUF) + }); err != nil { + return 0, err + } + return size, serr +} + +func inspectWriteBuffer(c syscall.RawConn) (int, error) { + var size int + var serr error + if err := c.Control(func(fd uintptr) { + size, serr = windows.GetsockoptInt(windows.Handle(fd), windows.SOL_SOCKET, windows.SO_SNDBUF) + }); err != nil { + return 0, err + } + return size, serr +} + +type packetInfo struct { + addr netip.Addr +} + +func (i *packetInfo) OOB() []byte { return nil } diff --git a/third_party/quic-go/sys_conn_windows_test.go b/third_party/quic-go/sys_conn_windows_test.go new file mode 100644 index 0000000..9376627 --- /dev/null +++ b/third_party/quic-go/sys_conn_windows_test.go @@ -0,0 +1,30 @@ +//go:build windows + +package quic + +import ( + "net" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestWindowsConn(t *testing.T) { + t.Run("IPv4", func(t *testing.T) { + udpConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + require.NoError(t, err) + conn, err := newConn(udpConn, true) + require.NoError(t, err) + require.NoError(t, conn.Close()) + require.True(t, conn.capabilities().DF) + }) + + t.Run("IPv6", func(t *testing.T) { + udpConn, err := net.ListenUDP("udp6", &net.UDPAddr{IP: net.IPv6loopback, Port: 0}) + require.NoError(t, err) + conn, err := newConn(udpConn, false) + require.NoError(t, err) + require.NoError(t, conn.Close()) + require.False(t, conn.capabilities().DF) + }) +} diff --git a/third_party/quic-go/testutils/events/event_recorder.go b/third_party/quic-go/testutils/events/event_recorder.go new file mode 100644 index 0000000..9d0449c --- /dev/null +++ b/third_party/quic-go/testutils/events/event_recorder.go @@ -0,0 +1,94 @@ +package events + +import ( + "reflect" + "slices" + "sync" + "time" + + "github.com/apernet/quic-go/qlogwriter" +) + +// Event is a recorded event with the event time. +type Event struct { + Time time.Time + Event qlogwriter.Event +} + +// Trace is a qlog.Trace that returns a qlog recorder. +type Trace struct { + Recorder qlogwriter.Recorder +} + +var _ qlogwriter.Trace = &Trace{} + +func (t *Trace) AddProducer() qlogwriter.Recorder { + return t.Recorder +} + +func (t *Trace) SupportsSchemas(string) bool { + return true +} + +// Recorder is a qlog.Recorder that records events. +// Events can be retrieved using the Events method. +type Recorder struct { + mx sync.Mutex + events []Event +} + +var _ qlogwriter.Recorder = &Recorder{} + +// Events returns all recorded events. +// If filter is provided, only events of the given type(s) are returned. +func (r *Recorder) RecordEvent(ev qlogwriter.Event) { + r.mx.Lock() + r.events = append(r.events, Event{Time: time.Now(), Event: ev}) + r.mx.Unlock() +} + +// Events returns all recorded events, including the event time. +// If filter is provided, only events of the given type(s) are returned. +func (r *Recorder) Events(filter ...qlogwriter.Event) []qlogwriter.Event { + eventsWithTime := r.EventsWithTime(filter...) + events := make([]qlogwriter.Event, 0, len(eventsWithTime)) + for _, ev := range eventsWithTime { + events = append(events, ev.Event) + } + return events +} + +func (r *Recorder) EventsWithTime(filter ...qlogwriter.Event) []Event { + r.mx.Lock() + events := r.events + r.mx.Unlock() + + if len(filter) == 0 { + return events + } + + // Some events have the same name when serialized, but use different structs. + // We therefore need to filter by type, and can't use the event name. + filterTypes := make([]reflect.Type, 0, len(filter)) + for _, f := range filter { + filterTypes = append(filterTypes, reflect.TypeOf(f)) + } + + var filtered []Event + for _, ev := range events { + eventType := reflect.TypeOf(ev.Event) + if slices.Contains(filterTypes, eventType) { + filtered = append(filtered, ev) + } + } + return filtered +} + +// Clear clears the recorded events. +func (r *Recorder) Clear() { + r.mx.Lock() + r.events = nil + r.mx.Unlock() +} + +func (r *Recorder) Close() error { return nil } diff --git a/third_party/quic-go/testutils/events/event_recorder_test.go b/third_party/quic-go/testutils/events/event_recorder_test.go new file mode 100644 index 0000000..f14ca0d --- /dev/null +++ b/third_party/quic-go/testutils/events/event_recorder_test.go @@ -0,0 +1,101 @@ +package events + +import ( + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/stretchr/testify/require" +) + +func TestRecorder(t *testing.T) { + recorder := &Recorder{} + defer recorder.Close() + + recorder.RecordEvent(qlog.MTUUpdated{Value: 1000}) + recorder.RecordEvent(qlog.ALPNInformation{ChosenALPN: "foobar"}) + recorder.RecordEvent(qlog.ECNStateUpdated{State: qlog.ECNStateCapable}) + recorder.RecordEvent(qlog.MTUUpdated{Value: 1200}) + + require.Equal(t, + []qlogwriter.Event{ + qlog.MTUUpdated{Value: 1000}, + qlog.ALPNInformation{ChosenALPN: "foobar"}, + qlog.ECNStateUpdated{State: qlog.ECNStateCapable}, + qlog.MTUUpdated{Value: 1200}, + }, + recorder.Events(), + ) + + require.Empty(t, recorder.Events(qlog.PacketBuffered{})) + require.Equal(t, + []qlogwriter.Event{ + qlog.MTUUpdated{Value: 1000}, + qlog.MTUUpdated{Value: 1200}, + }, + recorder.Events(qlog.MTUUpdated{}), + ) + + recorder.Clear() + require.Empty(t, recorder.Events()) + require.Empty(t, recorder.Events(qlog.MTUUpdated{})) +} + +func TestRecorderFilterEventsSameName(t *testing.T) { + // some events have the same name when serialized, but use different structs + require.Equal(t, + qlog.PacketReceived{}.Name(), + qlog.VersionNegotiationReceived{}.Name(), + ) + + recorder := &Recorder{} + defer recorder.Close() + + recorder.RecordEvent(qlog.PacketReceived{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeHandshake}, + }) + recorder.RecordEvent(qlog.VersionNegotiationReceived{ + Header: qlog.PacketHeaderVersionNegotiation{}, + SupportedVersions: []qlog.Version{0xdeadbeef, 0xdecafbad}, + }) + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketReceived{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeHandshake}, + }, + }, + recorder.Events(qlog.PacketReceived{}), + ) + require.Equal(t, + []qlogwriter.Event{ + qlog.VersionNegotiationReceived{ + Header: qlog.PacketHeaderVersionNegotiation{}, + SupportedVersions: []qlog.Version{0xdeadbeef, 0xdecafbad}, + }, + }, + recorder.Events(qlog.VersionNegotiationReceived{}), + ) +} + +func TestRecorderEventsWithTime(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + recorder := &Recorder{} + start := time.Now() + recorder.RecordEvent(qlog.MTUUpdated{Value: 1000}) + time.Sleep(time.Minute) + recorder.RecordEvent(qlog.ECNStateUpdated{State: qlog.ECNStateCapable}) + time.Sleep(time.Minute) + recorder.RecordEvent(qlog.MTUUpdated{Value: 1200}) + + require.Equal(t, + []Event{ + {Time: start, Event: qlog.MTUUpdated{Value: 1000}}, + {Time: start.Add(2 * time.Minute), Event: qlog.MTUUpdated{Value: 1200}}, + }, + recorder.EventsWithTime(qlog.MTUUpdated{}), + ) + }) +} diff --git a/third_party/quic-go/testutils/frames.go b/third_party/quic-go/testutils/frames.go new file mode 100644 index 0000000..547e0b0 --- /dev/null +++ b/third_party/quic-go/testutils/frames.go @@ -0,0 +1,26 @@ +package testutils + +import "github.com/apernet/quic-go/internal/wire" + +type ( + Frame = wire.Frame + AckFrame = wire.AckFrame + ConnectionCloseFrame = wire.ConnectionCloseFrame + CryptoFrame = wire.CryptoFrame + DataBlockedFrame = wire.DataBlockedFrame + HandshakeDoneFrame = wire.HandshakeDoneFrame + MaxDataFrame = wire.MaxDataFrame + MaxStreamDataFrame = wire.MaxStreamDataFrame + MaxStreamsFrame = wire.MaxStreamsFrame + NewConnectionIDFrame = wire.NewConnectionIDFrame + NewTokenFrame = wire.NewTokenFrame + PathChallengeFrame = wire.PathChallengeFrame + PathResponseFrame = wire.PathResponseFrame + PingFrame = wire.PingFrame + ResetStreamFrame = wire.ResetStreamFrame + RetireConnectionIDFrame = wire.RetireConnectionIDFrame + StopSendingFrame = wire.StopSendingFrame + StreamDataBlockedFrame = wire.StreamDataBlockedFrame + StreamFrame = wire.StreamFrame + StreamsBlockedFrame = wire.StreamsBlockedFrame +) diff --git a/third_party/quic-go/testutils/simnet/README.md b/third_party/quic-go/testutils/simnet/README.md new file mode 100644 index 0000000..c7649a8 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/README.md @@ -0,0 +1,14 @@ +# simnet + +This package is based on @MarcoPolo's [simnet](https://github.com/marcopolo/simnet) package. + +A small Go library for simulating packet networks in-process. It provides +drop-in `net.PacketConn` endpoints connected through configurable virtual links +with latency and MTU constraints. Useful for testing networking code +without sockets or root privileges. + +- **Drop-in API**: implements `net.PacketConn` +- **Realistic links**: per-direction latency and MTU +- **Packet queuing**: priority queue for scheduled packet delivery +- **Routers**: perfect delivery, fixed-latency, simple firewall/NAT-like routing +- **Deterministic testing**: opt-in `synctest`-based tests for time control diff --git a/third_party/quic-go/testutils/simnet/queue.go b/third_party/quic-go/testutils/simnet/queue.go new file mode 100644 index 0000000..e23196d --- /dev/null +++ b/third_party/quic-go/testutils/simnet/queue.go @@ -0,0 +1,134 @@ +package simnet + +import ( + "container/heap" + "sync" + "time" +) + +// queue is a priority queue that delivers packets at their scheduled delivery time +type queue struct { + mu sync.Mutex + packets packetHeap + newPacket chan struct{} + closed bool + pushCount int +} + +func newQueue() *queue { + q := &queue{ + newPacket: make(chan struct{}, 1), + } + heap.Init(&q.packets) + return q +} + +// Enqueue adds a packet to the queue +func (q *queue) Enqueue(p *packetWithDeliveryTime) { + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return + } + q.pushCount++ + heap.Push(&q.packets, packetWithDeliveryTimeAndOrder{packetWithDeliveryTime: p, count: q.pushCount}) + + // Signal that a new packet arrived (non-blocking) + select { + case q.newPacket <- struct{}{}: + default: + } +} + +// Dequeue removes and returns the next packet when it's ready for delivery +// This blocks until a packet is available AND its delivery time has been reached +// Uses a timer that can be reset if a packet with earlier delivery time arrives +func (q *queue) Dequeue() (*packetWithDeliveryTime, bool) { + timer := time.NewTimer(time.Hour) + timer.Stop() + + for { + q.mu.Lock() + + if q.closed { + q.mu.Unlock() + timer.Stop() + return nil, false + } + + if len(q.packets) == 0 { + // no packets, wait for one to arrive + q.mu.Unlock() + <-q.newPacket + timer.Stop() + continue + } + + earliest := q.packets[0] + earliestTime := earliest.DeliveryTime + + now := time.Now() + if now.Before(earliestTime) { + // not ready yet, wait until delivery time or new packet + waitDuration := earliestTime.Sub(now) + timer.Reset(waitDuration) + q.mu.Unlock() + + select { + case <-timer.C: + continue + case <-q.newPacket: + // new packet arrived, might have earlier delivery time + timer.Stop() + continue + } + } + + // Packet is ready, remove from queue and return it + po := heap.Pop(&q.packets).(packetWithDeliveryTimeAndOrder) + p := po.packetWithDeliveryTime + + q.mu.Unlock() + + return p, true + } +} + +// Close closes the queue +func (q *queue) Close() { + q.mu.Lock() + defer q.mu.Unlock() + + q.closed = true + close(q.newPacket) +} + +type packetWithDeliveryTimeAndOrder struct { + count int + *packetWithDeliveryTime +} + +// packetHeap implements heap.Interface ordered by packet delivery time. +type packetHeap []packetWithDeliveryTimeAndOrder + +func (h packetHeap) Len() int { return len(h) } + +func (h packetHeap) Less(i, j int) bool { + return (h[i].DeliveryTime.Before(h[j].DeliveryTime) || h[i].DeliveryTime.Equal(h[j].DeliveryTime) && h[i].count < h[j].count) +} + +func (h packetHeap) Swap(i, j int) { + h[i], h[j] = h[j], h[i] +} + +func (h *packetHeap) Push(x any) { + *h = append(*h, x.(packetWithDeliveryTimeAndOrder)) +} + +func (h *packetHeap) Pop() any { + old := *h + n := len(old) + item := old[n-1] + *h = old[:n-1] + return item +} diff --git a/third_party/quic-go/testutils/simnet/queue_test.go b/third_party/quic-go/testutils/simnet/queue_test.go new file mode 100644 index 0000000..ece78e2 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/queue_test.go @@ -0,0 +1,83 @@ +package simnet + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestQueue(t *testing.T) { + q := newQueue() + baseTime := time.Now() + + // Enqueue 5 packets with different delivery times + // Two packets scheduled for the same time (t2) + p1 := &packetWithDeliveryTime{ + Packet: Packet{Data: []byte("packet1")}, + DeliveryTime: baseTime.Add(10 * time.Millisecond), + } + p2 := &packetWithDeliveryTime{ + Packet: Packet{Data: []byte("packet2")}, + DeliveryTime: baseTime.Add(20 * time.Millisecond), + } + p3 := &packetWithDeliveryTime{ + Packet: Packet{Data: []byte("packet3")}, + DeliveryTime: baseTime.Add(20 * time.Millisecond), // Same time as p2 + } + p4 := &packetWithDeliveryTime{ + Packet: Packet{Data: []byte("packet4")}, + DeliveryTime: baseTime.Add(30 * time.Millisecond), + } + p5 := &packetWithDeliveryTime{ + Packet: Packet{Data: []byte("packet5")}, + DeliveryTime: baseTime.Add(5 * time.Millisecond), + } + + // Enqueue in non-chronological order + q.Enqueue(p1) + q.Enqueue(p2) + q.Enqueue(p3) + q.Enqueue(p4) + q.Enqueue(p5) + + // Dequeue should return packets in order: p5, p1, p2, p3, p4 + // p2 and p3 have same time, but p2 was enqueued first + received, ok := q.Dequeue() + require.True(t, ok) + require.Equal(t, "packet5", string(received.Data)) + + received, ok = q.Dequeue() + require.True(t, ok) + require.Equal(t, "packet1", string(received.Data)) + + received, ok = q.Dequeue() + require.True(t, ok) + require.Equal(t, "packet2", string(received.Data)) + + received, ok = q.Dequeue() + require.True(t, ok) + require.Equal(t, "packet3", string(received.Data)) + + received, ok = q.Dequeue() + require.True(t, ok) + require.Equal(t, "packet4", string(received.Data)) +} + +func TestQueueClose(t *testing.T) { + q := newQueue() + q.Close() + + _, ok := q.Dequeue() + require.False(t, ok) + + // enqueue after close should be ignored + p := &packetWithDeliveryTime{ + Packet: Packet{Data: []byte("packet")}, + DeliveryTime: time.Now(), + } + q.Enqueue(p) + // dequeue should still return false + _, ok = q.Dequeue() + require.False(t, ok) +} diff --git a/third_party/quic-go/testutils/simnet/router.go b/third_party/quic-go/testutils/simnet/router.go new file mode 100644 index 0000000..a8034d0 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/router.go @@ -0,0 +1,149 @@ +package simnet + +import ( + "errors" + "net" + "net/netip" + "sync" + "time" +) + +type ipPortKey struct { + ip string + port uint16 + isUDP bool +} + +func (k *ipPortKey) FromNetAddr(addr net.Addr) error { + switch addr := addr.(type) { + case *net.UDPAddr: + *k = ipPortKey{ + ip: string(addr.IP), + port: uint16(addr.Port), + isUDP: true, + } + return nil + case *net.TCPAddr: + *k = ipPortKey{ + ip: string(addr.IP), + port: uint16(addr.Port), + isUDP: false, + } + return nil + default: + ip, err := netip.ParseAddrPort(addr.String()) + if err != nil { + return err + } + *k = ipPortKey{ + ip: string(ip.Addr().AsSlice()), + port: ip.Port(), + isUDP: addr.Network() == "udp", + } + return nil + } +} + +type addrMap[V any] struct { + mu sync.Mutex + nodes map[ipPortKey]V +} + +func (m *addrMap[V]) Get(addr net.Addr) (V, bool) { + m.mu.Lock() + defer m.mu.Unlock() + var v V + if len(m.nodes) == 0 { + return v, false + } + var k ipPortKey + if err := k.FromNetAddr(addr); err != nil { + return v, false + } + v, ok := m.nodes[k] + return v, ok +} + +func (m *addrMap[V]) Set(addr net.Addr, v V) error { + m.mu.Lock() + defer m.mu.Unlock() + if m.nodes == nil { + m.nodes = make(map[ipPortKey]V) + } + + var k ipPortKey + if err := k.FromNetAddr(addr); err != nil { + return err + } + m.nodes[k] = v + return nil +} + +func (m *addrMap[V]) Delete(addr net.Addr) error { + m.mu.Lock() + defer m.mu.Unlock() + if m.nodes == nil { + m.nodes = make(map[ipPortKey]V) + } + + var k ipPortKey + if err := k.FromNetAddr(addr); err != nil { + return err + } + delete(m.nodes, k) + return nil +} + +// PerfectRouter is a router that has no latency or jitter and can route to +// every node +type PerfectRouter struct { + nodes addrMap[PacketReceiver] +} + +// SendPacket implements Router. +func (r *PerfectRouter) SendPacket(p Packet) error { + conn, ok := r.nodes.Get(p.To) + if !ok { + return errors.New("unknown destination") + } + + conn.RecvPacket(p) + return nil +} + +func (r *PerfectRouter) AddNode(addr net.Addr, conn PacketReceiver) { + r.nodes.Set(addr, conn) +} + +func (r *PerfectRouter) RemoveNode(addr net.Addr) { + r.nodes.Delete(addr) +} + +var _ Router = &PerfectRouter{} + +type DelayedPacketReceiver struct { + inner PacketReceiver + delay time.Duration +} + +func (r *DelayedPacketReceiver) RecvPacket(p Packet) { + time.AfterFunc(r.delay, func() { r.inner.RecvPacket(p) }) +} + +type FixedLatencyRouter struct { + PerfectRouter + latency time.Duration +} + +func (r *FixedLatencyRouter) SendPacket(p Packet) error { + return r.PerfectRouter.SendPacket(p) +} + +func (r *FixedLatencyRouter) AddNode(addr net.Addr, conn PacketReceiver) { + r.PerfectRouter.AddNode(addr, &DelayedPacketReceiver{ + inner: conn, + delay: r.latency, + }) +} + +var _ Router = &FixedLatencyRouter{} diff --git a/third_party/quic-go/testutils/simnet/simconn.go b/third_party/quic-go/testutils/simnet/simconn.go new file mode 100644 index 0000000..9144a25 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/simconn.go @@ -0,0 +1,250 @@ +package simnet + +import ( + "errors" + "net" + "slices" + "sync" + "sync/atomic" + "time" +) + +var ErrDeadlineExceeded = errors.New("deadline exceeded") + +type PacketReceiver interface { + RecvPacket(p Packet) +} + +// Router handles routing of packets between simulated connections. +// Implementations are responsible for delivering packets to their destinations. +type Router interface { + SendPacket(p Packet) error + AddNode(addr net.Addr, receiver PacketReceiver) +} + +type Packet struct { + To net.Addr + From net.Addr + + Data []byte +} + +// SimConn is a simulated network connection that implements net.PacketConn. +// It provides packet-based communication through a Router for testing and +// simulation purposes. All send/recv operations are handled through the +// Router's packet delivery mechanism. +type SimConn struct { + mu sync.Mutex + closed bool + closedChan chan struct{} + deadlineUpdated chan struct{} + + packetsSent atomic.Uint64 + packetsRcvd atomic.Uint64 + bytesSent atomic.Int64 + bytesRcvd atomic.Int64 + + router Router + + myAddr *net.UDPAddr + myLocalAddr net.Addr + packetsToRead chan Packet + + // Controls whether to block when receiving packets if our buffer is full. + // If false, drops packets. + recvBackPressure bool + + readDeadline time.Time + writeDeadline time.Time +} + +var _ net.PacketConn = &SimConn{} + +// NewSimConn creates a new simulated connection that drops packets if the +// receive buffer is full. +func NewSimConn(addr *net.UDPAddr, rtr Router) *SimConn { + return newSimConn(addr, rtr, false) +} + +// NewBlockingSimConn creates a new simulated connection that blocks if the +// receive buffer is full. Does not drop packets. +func NewBlockingSimConn(addr *net.UDPAddr, rtr Router) *SimConn { + return newSimConn(addr, rtr, true) +} + +func newSimConn(addr *net.UDPAddr, rtr Router, block bool) *SimConn { + c := &SimConn{ + recvBackPressure: block, + router: rtr, + myAddr: addr, + packetsToRead: make(chan Packet, 32), + closedChan: make(chan struct{}), + deadlineUpdated: make(chan struct{}, 1), + } + rtr.AddNode(addr, c) + return c +} + +type ConnStats struct { + BytesSent int + BytesRcvd int + PacketsSent int + PacketsRcvd int +} + +func (c *SimConn) Stats() ConnStats { + return ConnStats{ + BytesSent: int(c.bytesSent.Load()), + BytesRcvd: int(c.bytesRcvd.Load()), + PacketsSent: int(c.packetsSent.Load()), + PacketsRcvd: int(c.packetsRcvd.Load()), + } +} + +// SetReadBuffer only exists to quell the warning message from quic-go +func (c *SimConn) SetReadBuffer(n int) error { + return nil +} + +// SetWriteBuffer only exists to quell the warning message from quic-go +func (c *SimConn) SetWriteBuffer(n int) error { + return nil +} + +func (c *SimConn) RecvPacket(p Packet) { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return + } + c.mu.Unlock() + c.packetsRcvd.Add(1) + c.bytesRcvd.Add(int64(len(p.Data))) + + if c.recvBackPressure { + select { + case c.packetsToRead <- p: + case <-c.closedChan: + // if the connection is closed, drop the packet + return + } + } else { + select { + case c.packetsToRead <- p: + default: + // drop the packet if the channel is full + } + } +} + +func (c *SimConn) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return nil + } + c.closed = true + close(c.closedChan) + return nil +} + +func (c *SimConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return 0, nil, net.ErrClosed + } + deadline := c.readDeadline + c.mu.Unlock() + + if !deadline.IsZero() && !time.Now().Before(deadline) { + return 0, nil, ErrDeadlineExceeded + } + + var pkt Packet + var deadlineTimer <-chan time.Time + if !deadline.IsZero() { + deadlineTimer = time.After(time.Until(deadline)) + } + + select { + case pkt = <-c.packetsToRead: + case <-c.closedChan: + return 0, nil, net.ErrClosed + case <-c.deadlineUpdated: + return c.ReadFrom(p) + case <-deadlineTimer: + return 0, nil, ErrDeadlineExceeded + } + + n = copy(p, pkt.Data) + // if the provided buffer is not enough to read the whole packet, we drop + // the rest of the data. this is similar to what `recvfrom` does on Linux + // and macOS. + return n, pkt.From, nil +} + +func (c *SimConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return 0, net.ErrClosed + } + deadline := c.writeDeadline + c.mu.Unlock() + + if !deadline.IsZero() && !time.Now().Before(deadline) { + return 0, ErrDeadlineExceeded + } + + c.packetsSent.Add(1) + c.bytesSent.Add(int64(len(p))) + + pkt := Packet{ + From: c.myAddr, + To: addr, + Data: slices.Clone(p), + } + return len(p), c.router.SendPacket(pkt) +} + +func (c *SimConn) UnicastAddr() net.Addr { + return c.myAddr +} + +func (c *SimConn) LocalAddr() net.Addr { + if c.myLocalAddr != nil { + return c.myLocalAddr + } + return c.myAddr +} + +func (c *SimConn) SetDeadline(t time.Time) error { + c.mu.Lock() + defer c.mu.Unlock() + c.readDeadline = t + c.writeDeadline = t + select { + case c.deadlineUpdated <- struct{}{}: + default: + } + return nil +} + +func (c *SimConn) SetReadDeadline(t time.Time) error { + c.mu.Lock() + defer c.mu.Unlock() + c.readDeadline = t + select { + case c.deadlineUpdated <- struct{}{}: + default: + } + return nil +} + +func (c *SimConn) SetWriteDeadline(t time.Time) error { + c.mu.Lock() + defer c.mu.Unlock() + c.writeDeadline = t + return nil +} diff --git a/third_party/quic-go/testutils/simnet/simconn_test.go b/third_party/quic-go/testutils/simnet/simconn_test.go new file mode 100644 index 0000000..6886fb9 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/simconn_test.go @@ -0,0 +1,187 @@ +package simnet + +import ( + "crypto/rand" + "net" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func randomPublicIPv4() net.IP { +start: + ip := make([]byte, 4) + rand.Read(ip[:]) + if net.IP(ip).IsPrivate() || net.IP(ip).IsLoopback() || net.IP(ip).IsLinkLocalUnicast() { + goto start + } + return ip +} + +func TestSimConnBasicConnectivity(t *testing.T) { + router := &PerfectRouter{} + + // Create two endpoints + addr1 := &net.UDPAddr{IP: randomPublicIPv4(), Port: 1234} + addr2 := &net.UDPAddr{IP: randomPublicIPv4(), Port: 1234} + + conn1 := NewSimConn(addr1, router) + conn2 := NewSimConn(addr2, router) + + // Test sending data from conn1 to conn2 + testData := []byte("hello world") + n, err := conn1.WriteTo(testData, addr2) + require.NoError(t, err) + require.Equal(t, len(testData), n) + + // Read data from conn2 + buf := make([]byte, 1024) + n, addr, err := conn2.ReadFrom(buf) + require.NoError(t, err) + require.Equal(t, testData, buf[:n]) + require.Equal(t, addr1, addr) + + // Check stats + stats1 := conn1.Stats() + require.Equal(t, len(testData), stats1.BytesSent) + require.Equal(t, 1, stats1.PacketsSent) + + stats2 := conn2.Stats() + require.Equal(t, len(testData), stats2.BytesRcvd) + require.Equal(t, 1, stats2.PacketsRcvd) +} + +func TestSimConnDeadlines(t *testing.T) { + router := &PerfectRouter{} + + addr1 := &net.UDPAddr{IP: randomPublicIPv4(), Port: 1234} + conn := NewSimConn(addr1, router) + + t.Run("read deadline", func(t *testing.T) { + deadline := time.Now().Add(10 * time.Millisecond) + err := conn.SetReadDeadline(deadline) + require.NoError(t, err) + + buf := make([]byte, 1024) + _, _, err = conn.ReadFrom(buf) + require.ErrorIs(t, err, ErrDeadlineExceeded) + }) + + t.Run("write deadline", func(t *testing.T) { + deadline := time.Now().Add(-time.Second) // Already expired + err := conn.SetWriteDeadline(deadline) + require.NoError(t, err) + + _, err = conn.WriteTo([]byte("test"), &net.UDPAddr{}) + require.ErrorIs(t, err, ErrDeadlineExceeded) + }) +} + +func TestSimConnClose(t *testing.T) { + router := &PerfectRouter{} + + addr1 := &net.UDPAddr{IP: randomPublicIPv4(), Port: 1234} + conn := NewSimConn(addr1, router) + + err := conn.Close() + require.NoError(t, err) + + // Verify operations fail after close + _, err = conn.WriteTo([]byte("test"), addr1) + require.ErrorIs(t, err, net.ErrClosed) + + buf := make([]byte, 1024) + _, _, err = conn.ReadFrom(buf) + require.ErrorIs(t, err, net.ErrClosed) + + // Second close should not error + err = conn.Close() + require.NoError(t, err) +} + +func TestSimConnDeadlinesWithLatency(t *testing.T) { + router := &FixedLatencyRouter{ + PerfectRouter: PerfectRouter{}, + latency: 100 * time.Millisecond, + } + + addr1 := &net.UDPAddr{IP: randomPublicIPv4(), Port: 1234} + addr2 := &net.UDPAddr{IP: randomPublicIPv4(), Port: 1234} + + conn1 := NewSimConn(addr1, router) + conn2 := NewSimConn(addr2, router) + + reset := func() { + router.RemoveNode(addr1) + router.RemoveNode(addr2) + + conn1 = NewSimConn(addr1, router) + conn2 = NewSimConn(addr2, router) + } + + t.Run("write succeeds within deadline", func(t *testing.T) { + deadline := time.Now().Add(200 * time.Millisecond) + err := conn1.SetWriteDeadline(deadline) + require.NoError(t, err) + + n, err := conn1.WriteTo([]byte("test"), addr2) + require.NoError(t, err) + require.Equal(t, 4, n) + reset() + }) + + t.Run("write fails after past deadline", func(t *testing.T) { + deadline := time.Now().Add(-time.Second) // Already expired + err := conn1.SetWriteDeadline(deadline) + require.NoError(t, err) + + _, err = conn1.WriteTo([]byte("test"), addr2) + require.ErrorIs(t, err, ErrDeadlineExceeded) + reset() + }) + + t.Run("read succeeds within deadline", func(t *testing.T) { + // Reset deadline and send a message + conn2.SetReadDeadline(time.Time{}) + testData := []byte("hello") + deadline := time.Now().Add(200 * time.Millisecond) + conn1.SetWriteDeadline(deadline) + _, err := conn1.WriteTo(testData, addr2) + require.NoError(t, err) + + // Set read deadline and try to read + deadline = time.Now().Add(200 * time.Millisecond) + err = conn2.SetReadDeadline(deadline) + require.NoError(t, err) + + buf := make([]byte, 1024) + n, addr, err := conn2.ReadFrom(buf) + require.NoError(t, err) + require.Equal(t, addr1, addr) + require.Equal(t, testData, buf[:n]) + reset() + }) + + t.Run("read fails after deadline", func(t *testing.T) { + defer reset() + // Set a short deadline + deadline := time.Now().Add(50 * time.Millisecond) // Less than router latency + err := conn2.SetReadDeadline(deadline) + require.NoError(t, err) + + var wg sync.WaitGroup + defer wg.Wait() + wg.Go(func() { + // Send data after setting deadline + _, err := conn1.WriteTo([]byte("test"), addr2) + require.NoError(t, err) + }) + + // Read should fail due to deadline + buf := make([]byte, 1024) + _, _, err = conn2.ReadFrom(buf) + require.ErrorIs(t, err, ErrDeadlineExceeded) + }) +} diff --git a/third_party/quic-go/testutils/simnet/simlink.go b/third_party/quic-go/testutils/simnet/simlink.go new file mode 100644 index 0000000..3730d95 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/simlink.go @@ -0,0 +1,145 @@ +package simnet + +import ( + "net" + "sync" + "time" +) + +// packetWithDeliveryTime holds a packet along with its scheduled delivery time +type packetWithDeliveryTime struct { + Packet + DeliveryTime time.Time +} + +// LinkSettings defines the network characteristics for a simulated link direction +type LinkSettings struct { + // MTU (Maximum Transmission Unit) specifies the maximum packet size in bytes + MTU int +} + +// SimulatedLink simulates a bidirectional network link with variable latency and MTU constraints +type SimulatedLink struct { + // Internal state for lifecycle management + wg sync.WaitGroup + + // Queues for packet delivery timing + downstreamQueue *queue + upstreamQueue *queue + + // Configuration for link characteristics + UplinkSettings LinkSettings + DownlinkSettings LinkSettings + + // Latency specifies a fixed network delay for downlink packets + // If both Latency and LatencyFunc are set, LatencyFunc takes precedence + Latency time.Duration + + // LatencyFunc computes the network delay for each downlink packet + // This allows variable latency based on packet source/destination + // If nil, Latency field is used instead + LatencyFunc func(Packet) time.Duration + + // Packet routing interfaces + UploadPacket Router + downloadPacket PacketReceiver +} + +func (l *SimulatedLink) AddNode(addr net.Addr, receiver PacketReceiver) { + l.downloadPacket = receiver +} + +func (l *SimulatedLink) Start() { + if l.downloadPacket == nil { + panic("SimulatedLink.Start() called without having added a packet receiver") + } + + // Sane defaults + if l.DownlinkSettings.MTU == 0 { + l.DownlinkSettings.MTU = 1400 + } + if l.UplinkSettings.MTU == 0 { + l.UplinkSettings.MTU = 1400 + } + + l.downstreamQueue = newQueue() + l.upstreamQueue = newQueue() + + l.wg.Go(func() { l.backgroundDownlink() }) + l.wg.Go(func() { l.backgroundUplink() }) +} + +func (l *SimulatedLink) Close() error { + l.downstreamQueue.Close() + l.upstreamQueue.Close() + l.wg.Wait() + return nil +} + +func (l *SimulatedLink) backgroundDownlink() { + for { + // Dequeue a packet (this will block until packet is ready for delivery) + // Dequeue() returns false when the queue is closed + p, ok := l.downstreamQueue.Dequeue() + if !ok { + return + } + + // Deliver the packet + l.downloadPacket.RecvPacket(p.Packet) + } +} + +func (l *SimulatedLink) backgroundUplink() { + for { + // Dequeue a packet (this will block until packet is ready for delivery) + // Dequeue() returns false when the queue is closed + p, ok := l.upstreamQueue.Dequeue() + if !ok { + return + } + + // Deliver the packet + _ = l.UploadPacket.SendPacket(p.Packet) + } +} + +func (l *SimulatedLink) SendPacket(p Packet) error { + if len(p.Data) > l.UplinkSettings.MTU { + // Drop packet if it's too large + return nil + } + + // Uplink has no latency - packets are delivered immediately + deliveryTime := time.Now() + + // Enqueue packet with delivery time + l.upstreamQueue.Enqueue(&packetWithDeliveryTime{ + Packet: p, + DeliveryTime: deliveryTime, + }) + + return nil +} + +func (l *SimulatedLink) RecvPacket(p Packet) { + if len(p.Data) > l.DownlinkSettings.MTU { + // Drop packet if it's too large + return + } + + // Calculate delivery time based on downlink latency + var latency time.Duration + if l.LatencyFunc != nil { + latency = l.LatencyFunc(p) + } else { + latency = l.Latency + } + deliveryTime := time.Now().Add(latency) + + // Enqueue packet with delivery time + l.downstreamQueue.Enqueue(&packetWithDeliveryTime{ + Packet: p, + DeliveryTime: deliveryTime, + }) +} diff --git a/third_party/quic-go/testutils/simnet/simlink_test.go b/third_party/quic-go/testutils/simnet/simlink_test.go new file mode 100644 index 0000000..f44bb0d --- /dev/null +++ b/third_party/quic-go/testutils/simnet/simlink_test.go @@ -0,0 +1,153 @@ +package simnet + +import ( + "fmt" + "math" + "net" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" +) + +type testRouter struct { + onSend func(p Packet) + onRecv func(p Packet) +} + +func (r *testRouter) SendPacket(p Packet) error { + r.onSend(p) + return nil +} + +func (r *testRouter) RecvPacket(p Packet) { + r.onRecv(p) +} + +func (r *testRouter) AddNode(addr net.Addr, receiver PacketReceiver) { + r.onRecv = receiver.RecvPacket +} + +func TestLatency(t *testing.T) { + for _, testUpload := range []bool{true, false} { + t.Run(fmt.Sprintf("testing upload=%t", testUpload), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const downlinkLatency = 10 * time.Millisecond + const MTU = 1400 + linkSettings := LinkSettings{ + MTU: MTU, + } + + recvStartTimeChan := make(chan time.Time, 1) + recvStarted := false + packetHandler := func(p Packet) { + if !recvStarted { + recvStarted = true + recvStartTimeChan <- time.Now() + } + } + + router := &testRouter{} + if testUpload { + router.onSend = packetHandler + } else { + router.onRecv = packetHandler + } + link := SimulatedLink{ + UplinkSettings: linkSettings, + DownlinkSettings: linkSettings, + LatencyFunc: func(p Packet) time.Duration { return downlinkLatency }, + UploadPacket: router, + downloadPacket: router, + } + + link.Start() + + chunk := make([]byte, MTU) + sendStartTime := time.Now() + if testUpload { + _ = link.SendPacket(Packet{Data: chunk}) + } else { + link.RecvPacket(Packet{Data: chunk}) + } + + // Wait for delayed packets to be sent + time.Sleep(40 * time.Millisecond) + + link.Close() + recvStartTime := <-recvStartTimeChan + + observedLatency := recvStartTime.Sub(sendStartTime) + // Uplink is now instant (no latency), only downlink has latency + var expectedLatency time.Duration + if testUpload { + // Uplink test: expect near-zero latency + expectedLatency = 0 + t.Logf("observed latency: %s (uplink is instant)", observedLatency) + if observedLatency > 5*time.Millisecond { + t.Fatalf("observed latency %s is too high for instant uplink", observedLatency) + } + } else { + // Downlink test: expect configured latency + expectedLatency = downlinkLatency + percentErrorLatency := math.Abs(observedLatency.Seconds()-expectedLatency.Seconds()) / expectedLatency.Seconds() + t.Logf("observed latency: %s, expected latency: %s, percent error: %f", observedLatency, expectedLatency, percentErrorLatency) + if percentErrorLatency > 0.20 { + t.Fatalf("observed latency %s is wrong", observedLatency) + } + } + }) + }) + } +} + +func TestMTUEnforcement(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const MTU = 1400 + linkSettings := LinkSettings{ + MTU: MTU, + } + + var packetsReceived atomic.Uint32 + router := &testRouter{ + onSend: func(p Packet) { packetsReceived.Add(1) }, + onRecv: func(p Packet) { packetsReceived.Add(1) }, + } + link := SimulatedLink{ + UplinkSettings: linkSettings, + DownlinkSettings: linkSettings, + UploadPacket: router, + downloadPacket: router, + } + + link.Start() + + // Send a packet that fits within MTU - should be delivered + smallPacket := make([]byte, MTU) + err := link.SendPacket(Packet{Data: smallPacket}) + require.NoError(t, err) + + // Send a packet that exceeds MTU - should be dropped + largePacket := make([]byte, MTU+1) + err = link.SendPacket(Packet{Data: largePacket}) + require.NoError(t, err) // SendPacket returns nil even when dropping + + // Receive a packet that fits within MTU - should be delivered + link.RecvPacket(Packet{Data: smallPacket}) + + // Receive a packet that exceeds MTU - should be dropped + link.RecvPacket(Packet{Data: largePacket}) + + // Wait for packets to be processed + time.Sleep(10 * time.Millisecond) + + link.Close() + + // Only packets within MTU should be received (2 packets: 1 from SendPacket, 1 from RecvPacket) + if packetsReceived.Load() != 2 { + t.Fatalf("expected 2 packets to be received, got %d", packetsReceived.Load()) + } + }) +} diff --git a/third_party/quic-go/testutils/simnet/simnet.go b/third_party/quic-go/testutils/simnet/simnet.go new file mode 100644 index 0000000..adf6754 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/simnet.go @@ -0,0 +1,71 @@ +package simnet + +import ( + "errors" + "fmt" + "net" + "time" +) + +// Simnet is a simulated network that manages connections between nodes +// with configurable network conditions. +type Simnet struct { + Router Router + + links []*SimulatedLink +} + +// NodeBiDiLinkSettings defines the bidirectional link settings for a network node. +// It specifies separate configurations for downlink (incoming) and uplink (outgoing) +// traffic, allowing asymmetric network conditions to be simulated. +type NodeBiDiLinkSettings struct { + // Downlink configures the settings for incoming traffic to this node + Downlink LinkSettings + // Uplink configures the settings for outgoing traffic from this node + Uplink LinkSettings + + // Latency specifies a fixed network delay for downlink packets only + // If both Latency and LatencyFunc are set, LatencyFunc takes precedence + Latency time.Duration + + // LatencyFunc computes the network delay for each downlink packet + // This allows variable latency based on packet source/destination + // If nil, Latency field is used instead + LatencyFunc func(Packet) time.Duration +} + +func (n *Simnet) Start() error { + for _, link := range n.links { + link.Start() + } + return nil +} + +func (n *Simnet) Close() error { + var errs error + for _, link := range n.links { + err := link.Close() + if err != nil { + errs = errors.Join(errs, err) + } + } + if errs != nil { + return fmt.Errorf("failed to close some links: %w", errs) + } + return nil +} + +func (n *Simnet) NewEndpoint(addr *net.UDPAddr, linkSettings NodeBiDiLinkSettings) *SimConn { + link := &SimulatedLink{ + DownlinkSettings: linkSettings.Downlink, + UplinkSettings: linkSettings.Uplink, + Latency: linkSettings.Latency, + LatencyFunc: linkSettings.LatencyFunc, + UploadPacket: n.Router, + } + c := NewBlockingSimConn(addr, link) + + n.links = append(n.links, link) + n.Router.AddNode(addr, link) + return c +} diff --git a/third_party/quic-go/testutils/simnet/simnet_synctest_test.go b/third_party/quic-go/testutils/simnet/simnet_synctest_test.go new file mode 100644 index 0000000..5848222 --- /dev/null +++ b/third_party/quic-go/testutils/simnet/simnet_synctest_test.go @@ -0,0 +1,59 @@ +package simnet + +import ( + "math" + "net" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" +) + +func newConn(simnet *Simnet, address *net.UDPAddr, linkSettings NodeBiDiLinkSettings) *SimConn { + return simnet.NewEndpoint(address, linkSettings) +} + +func TestSimpleSimNet(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + router := &Simnet{Router: &PerfectRouter{}} + + const latency = 10 * time.Millisecond + linkSettings := NodeBiDiLinkSettings{ + Downlink: LinkSettings{}, + Uplink: LinkSettings{}, + Latency: latency, + } + + addressA := net.UDPAddr{ + IP: net.ParseIP("1.0.0.1"), + Port: 8000, + } + connA := newConn(router, &addressA, linkSettings) + addressB := net.UDPAddr{ + IP: net.ParseIP("1.0.0.2"), + Port: 8000, + } + connB := newConn(router, &addressB, linkSettings) + + router.Start() + defer router.Close() + + start := time.Now() + connA.WriteTo([]byte("hello"), &addressB) + buf := make([]byte, 1024) + n, from, err := connB.ReadFrom(buf) + require.NoError(t, err) + require.Equal(t, "hello", string(buf[:n])) + require.Equal(t, addressA.String(), from.String()) + observedLatency := time.Since(start) + + // Only downlink has latency now (uplink is instant) + expectedLatency := latency + percentDiff := math.Abs(float64(observedLatency-expectedLatency) / float64(expectedLatency)) + t.Logf("observed latency: %v, expected latency: %v, percent diff: %v", observedLatency, expectedLatency, percentDiff) + if percentDiff > 0.30 { + t.Fatalf("latency is wrong: %v. percent off: %v", observedLatency, percentDiff) + } + }) +} diff --git a/third_party/quic-go/testutils/testutils.go b/third_party/quic-go/testutils/testutils.go new file mode 100644 index 0000000..c14ea23 --- /dev/null +++ b/third_party/quic-go/testutils/testutils.go @@ -0,0 +1,108 @@ +// Package testutils contains utilities for simulating packet injection and man-in-the-middle (MITM) attacker tests. +// It is not supposed to be used for non-testing purposes. +// The API is not guaranteed to be stable. +package testutils + +import ( + "fmt" + + "github.com/apernet/quic-go/internal/handshake" + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/wire" +) + +// writePacket returns a new raw packet with the specified header and payload +func writePacket(hdr *wire.ExtendedHeader, data []byte) []byte { + b, err := hdr.Append(nil, hdr.Version) + if err != nil { + panic(fmt.Sprintf("failed to write header: %s", err)) + } + return append(b, data...) +} + +// packRawPayload returns a new raw payload containing given frames +func packRawPayload(version protocol.Version, frames []wire.Frame) []byte { + var b []byte + for _, cf := range frames { + var err error + b, err = cf.Append(b, version) + if err != nil { + panic(err) + } + } + return b +} + +// ComposeInitialPacket returns an Initial packet encrypted under key (the original destination connection ID) +// containing specified frames. +func ComposeInitialPacket( + srcConnID, destConnID, key protocol.ConnectionID, + token []byte, + frames []wire.Frame, + sentBy protocol.Perspective, + version protocol.Version, +) []byte { + sealer, _ := handshake.NewInitialAEAD(key, sentBy, version) + + // compose payload + var payload []byte + if len(frames) == 0 { + payload = make([]byte, protocol.MinInitialPacketSize) + } else { + payload = packRawPayload(version, frames) + } + + // compose Initial header + payloadSize := len(payload) + const pnLength = protocol.PacketNumberLen4 + length := payloadSize + int(pnLength) + sealer.Overhead() + hdr := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeInitial, + Token: token, + SrcConnectionID: srcConnID, + DestConnectionID: destConnID, + Length: protocol.ByteCount(length), + Version: version, + }, + PacketNumberLen: pnLength, + PacketNumber: 0x0, + } + + raw := writePacket(hdr, payload) + + // encrypt payload and header + payloadOffset := len(raw) - payloadSize + var encrypted []byte + encrypted = sealer.Seal(encrypted, payload, hdr.PacketNumber, raw[:payloadOffset]) + hdrBytes := raw[0:payloadOffset] + encrypted = append(hdrBytes, encrypted...) + pnOffset := payloadOffset - int(pnLength) // packet number offset + sealer.EncryptHeader( + encrypted[payloadOffset:payloadOffset+16], // first 16 bytes of payload (sample) + &encrypted[0], // first byte of header + encrypted[pnOffset:payloadOffset], // packet number bytes + ) + return encrypted +} + +// ComposeRetryPacket returns a new raw Retry Packet +func ComposeRetryPacket( + srcConnID protocol.ConnectionID, + destConnID protocol.ConnectionID, + origDestConnID protocol.ConnectionID, + token []byte, + version protocol.Version, +) []byte { + hdr := &wire.ExtendedHeader{ + Header: wire.Header{ + Type: protocol.PacketTypeRetry, + SrcConnectionID: srcConnID, + DestConnectionID: destConnID, + Token: token, + Version: version, + }, + } + data := writePacket(hdr, nil) + return append(data, handshake.GetRetryIntegrityTag(data, origDestConnID, version)[:]...) +} diff --git a/third_party/quic-go/token_store.go b/third_party/quic-go/token_store.go new file mode 100644 index 0000000..d93d305 --- /dev/null +++ b/third_party/quic-go/token_store.go @@ -0,0 +1,116 @@ +package quic + +import ( + "sync" + + list "github.com/apernet/quic-go/internal/utils/linkedlist" +) + +type singleOriginTokenStore struct { + tokens []*ClientToken + len int + p int +} + +func newSingleOriginTokenStore(size int) *singleOriginTokenStore { + return &singleOriginTokenStore{tokens: make([]*ClientToken, size)} +} + +func (s *singleOriginTokenStore) Add(token *ClientToken) { + s.tokens[s.p] = token + s.p = s.index(s.p + 1) + s.len = min(s.len+1, len(s.tokens)) +} + +func (s *singleOriginTokenStore) Pop() *ClientToken { + s.p = s.index(s.p - 1) + token := s.tokens[s.p] + s.tokens[s.p] = nil + s.len = max(s.len-1, 0) + return token +} + +func (s *singleOriginTokenStore) Len() int { + return s.len +} + +func (s *singleOriginTokenStore) index(i int) int { + mod := len(s.tokens) + return (i + mod) % mod +} + +type lruTokenStoreEntry struct { + key string + cache *singleOriginTokenStore +} + +type lruTokenStore struct { + mutex sync.Mutex + + m map[string]*list.Element[*lruTokenStoreEntry] + q *list.List[*lruTokenStoreEntry] + capacity int + singleOriginSize int +} + +var _ TokenStore = &lruTokenStore{} + +// NewLRUTokenStore creates a new LRU cache for tokens received by the client. +// maxOrigins specifies how many origins this cache is saving tokens for. +// tokensPerOrigin specifies the maximum number of tokens per origin. +func NewLRUTokenStore(maxOrigins, tokensPerOrigin int) TokenStore { + return &lruTokenStore{ + m: make(map[string]*list.Element[*lruTokenStoreEntry]), + q: list.New[*lruTokenStoreEntry](), + capacity: maxOrigins, + singleOriginSize: tokensPerOrigin, + } +} + +func (s *lruTokenStore) Put(key string, token *ClientToken) { + s.mutex.Lock() + defer s.mutex.Unlock() + + if el, ok := s.m[key]; ok { + entry := el.Value + entry.cache.Add(token) + s.q.MoveToFront(el) + return + } + + if s.q.Len() < s.capacity { + entry := &lruTokenStoreEntry{ + key: key, + cache: newSingleOriginTokenStore(s.singleOriginSize), + } + entry.cache.Add(token) + s.m[key] = s.q.PushFront(entry) + return + } + + elem := s.q.Back() + entry := elem.Value + delete(s.m, entry.key) + entry.key = key + entry.cache = newSingleOriginTokenStore(s.singleOriginSize) + entry.cache.Add(token) + s.q.MoveToFront(elem) + s.m[key] = elem +} + +func (s *lruTokenStore) Pop(key string) *ClientToken { + s.mutex.Lock() + defer s.mutex.Unlock() + + var token *ClientToken + if el, ok := s.m[key]; ok { + s.q.MoveToFront(el) + cache := el.Value.cache + token = cache.Pop() + if cache.Len() == 0 { + s.q.Remove(el) + delete(s.m, key) + } + } + return token +} diff --git a/third_party/quic-go/token_store_test.go b/third_party/quic-go/token_store_test.go new file mode 100644 index 0000000..0c7f90e --- /dev/null +++ b/third_party/quic-go/token_store_test.go @@ -0,0 +1,79 @@ +package quic + +import ( + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func mockToken(num int) *ClientToken { + return &ClientToken{data: fmt.Appendf(nil, "%d", num), rtt: 1337 * time.Millisecond} +} + +func TestTokenStoreSingleOrigin(t *testing.T) { + const origin = "localhost" + + s := NewLRUTokenStore(1, 3) + s.Put(origin, mockToken(1)) + s.Put(origin, mockToken(2)) + require.Equal(t, mockToken(2), s.Pop(origin)) + require.Equal(t, mockToken(1), s.Pop(origin)) + require.Nil(t, s.Pop(origin)) + + // now add more tokens than the cache size + s.Put(origin, mockToken(1)) + s.Put(origin, mockToken(2)) + s.Put(origin, mockToken(3)) + require.Equal(t, mockToken(3), s.Pop(origin)) + s.Put(origin, mockToken(4)) + s.Put(origin, mockToken(5)) + require.Equal(t, mockToken(5), s.Pop(origin)) + require.Equal(t, mockToken(4), s.Pop(origin)) + require.Equal(t, mockToken(2), s.Pop(origin)) + require.Nil(t, s.Pop(origin)) +} + +func TestTokenStoreMultipleOrigins(t *testing.T) { + s := NewLRUTokenStore(3, 4) + + s.Put("host1", mockToken(1)) + s.Put("host2", mockToken(2)) + s.Put("host3", mockToken(3)) + s.Put("host4", mockToken(4)) + require.Nil(t, s.Pop("host1")) + require.Equal(t, mockToken(2), s.Pop("host2")) + require.Equal(t, mockToken(3), s.Pop("host3")) + require.Equal(t, mockToken(4), s.Pop("host4")) +} + +func TestTokenStoreUpdates(t *testing.T) { + s := NewLRUTokenStore(3, 4) + s.Put("host1", mockToken(1)) + s.Put("host2", mockToken(2)) + s.Put("host3", mockToken(3)) + s.Put("host1", mockToken(11)) + // make sure one is evicted + s.Put("host4", mockToken(4)) + require.Nil(t, s.Pop("host2")) + require.Equal(t, mockToken(11), s.Pop("host1")) + require.Equal(t, mockToken(1), s.Pop("host1")) + require.Equal(t, mockToken(3), s.Pop("host3")) + require.Equal(t, mockToken(4), s.Pop("host4")) +} + +func TestTokenStoreEviction(t *testing.T) { + s := NewLRUTokenStore(3, 4) + + s.Put("host1", mockToken(1)) + s.Put("host2", mockToken(2)) + s.Put("host3", mockToken(3)) + require.Equal(t, mockToken(2), s.Pop("host2")) + require.Nil(t, s.Pop("host2")) + // host2 is now empty and should have been deleted, making space for host4 + s.Put("host4", mockToken(4)) + require.Equal(t, mockToken(1), s.Pop("host1")) + require.Equal(t, mockToken(3), s.Pop("host3")) + require.Equal(t, mockToken(4), s.Pop("host4")) +} diff --git a/third_party/quic-go/transport.go b/third_party/quic-go/transport.go new file mode 100644 index 0000000..bd08367 --- /dev/null +++ b/third_party/quic-go/transport.go @@ -0,0 +1,866 @@ +package quic + +import ( + "context" + "crypto/rand" + "crypto/tls" + "errors" + "fmt" + "net" + "sync" + "sync/atomic" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" +) + +// ErrTransportClosed is returned by the [Transport]'s Listen or Dial method after it was closed. +var ErrTransportClosed = &errTransportClosed{} + +type errTransportClosed struct { + err error +} + +func (e *errTransportClosed) Unwrap() []error { return []error{net.ErrClosed, e.err} } + +func (e *errTransportClosed) Error() string { + if e.err == nil { + return "quic: transport closed" + } + return fmt.Sprintf("quic: transport closed: %s", e.err) +} + +func (e *errTransportClosed) Is(target error) bool { + _, ok := target.(*errTransportClosed) + return ok +} + +var errListenerAlreadySet = errors.New("listener already set") + +type closePacket struct { + payload []byte + addr net.Addr + info packetInfo +} + +// The Transport is the central point to manage incoming and outgoing QUIC connections. +// QUIC demultiplexes connections based on their QUIC Connection IDs, not based on the 4-tuple. +// This means that a single UDP socket can be used for listening for incoming connections, as well as +// for dialing an arbitrary number of outgoing connections. +// A Transport handles a single net.PacketConn, and offers a range of configuration options +// compared to the simple helper functions like [Listen] and [Dial] that this package provides. +type Transport struct { + // A single net.PacketConn can only be handled by one Transport. + // Bad things will happen if passed to multiple Transports. + // + // A number of optimizations will be enabled if the connections implements the OOBCapablePacketConn interface, + // as a *net.UDPConn does. + // 1. It enables the Don't Fragment (DF) bit on the IP header. + // This is required to run DPLPMTUD (Path MTU Discovery, RFC 8899). + // 2. It enables reading of the ECN bits from the IP header. + // This allows the remote node to speed up its loss detection and recovery. + // 3. It uses batched syscalls (recvmmsg) to more efficiently receive packets from the socket. + // 4. It uses Generic Segmentation Offload (GSO) to efficiently send batches of packets (on Linux). + // + // After passing the connection to the Transport, it's invalid to call ReadFrom or WriteTo on the connection. + Conn net.PacketConn + + // The length of the connection ID in bytes. + // It can be any value between 1 and 20. + // Due to the increased risk of collisions, it is not recommended to use connection IDs shorter than 4 bytes. + // If unset, a 4 byte connection ID will be used. + ConnectionIDLength int + + // DisableGSO turns off UDP generic segmentation offload, at a cost in + // throughput. Set it when packets are rewritten after they leave the stack: + // the rewrite hits the combined packet and corrupts every segment but the + // first. The send still succeeds, so this cannot be detected automatically. + DisableGSO bool + + // Use for generating new connection IDs. + // This allows the application to control of the connection IDs used, + // which allows routing / load balancing based on connection IDs. + // All Connection IDs returned by the ConnectionIDGenerator MUST + // have the same length. + ConnectionIDGenerator ConnectionIDGenerator + + // The StatelessResetKey is used to generate stateless reset tokens. + // If no key is configured, sending of stateless resets is disabled. + // It is highly recommended to configure a stateless reset key, as stateless resets + // allow the peer to quickly recover from crashes and reboots of this node. + // See section 10.3 of RFC 9000 for details. + StatelessResetKey *StatelessResetKey + + // The TokenGeneratorKey is used to encrypt session resumption tokens. + // If no key is configured, a random key will be generated. + // If multiple servers are authoritative for the same domain, they should use the same key, + // see section 8.1.3 of RFC 9000 for details. + TokenGeneratorKey *TokenGeneratorKey + + // MaxTokenAge is the maximum age of the resumption token presented during the handshake. + // These tokens allow skipping address resumption when resuming a QUIC connection, + // and are especially useful when using 0-RTT. + // If not set, it defaults to 24 hours. + // See section 8.1.3 of RFC 9000 for details. + MaxTokenAge time.Duration + + // DisableVersionNegotiationPackets disables the sending of Version Negotiation packets. + // This can be useful if version information is exchanged out-of-band. + // It has no effect for clients. + DisableVersionNegotiationPackets bool + + // VerifySourceAddress decides if a connection attempt originating from unvalidated source + // addresses first needs to go through source address validation using QUIC's Retry mechanism, + // as described in RFC 9000 section 8.1.2. + // Note that the address passed to this callback is unvalidated, and might be spoofed in case + // of an attack. + // Validating the source address adds one additional network roundtrip to the handshake, + // and should therefore only be used if a suspiciously high number of incoming connection is recorded. + // For most use cases, wrapping the Allow function of a rate.Limiter will be a reasonable + // implementation of this callback (negating its return value). + VerifySourceAddress func(net.Addr) bool + + // ConnContext is called when the server accepts a new connection. To reject a connection return + // a non-nil error. + // The context is closed when the connection is closed, or when the handshake fails for any reason. + // The context returned from the callback is used to derive every other context used during the + // lifetime of the connection: + // * the context passed to crypto/tls (and used on the tls.ClientHelloInfo) + // * the context used in Config.QlogTrace + // * the context returned from Conn.Context + // * the context returned from SendStream.Context + // It is not used for dialed connections. + ConnContext func(context.Context, *ClientInfo) (context.Context, error) + + // A Tracer traces events that don't belong to a single QUIC connection. + // Recorder.Close is called when the transport is closed. + Tracer qlogwriter.Recorder + + mutex sync.Mutex + handlers map[protocol.ConnectionID]packetHandler + resetTokens map[protocol.StatelessResetToken]packetHandler + + initOnce sync.Once + initErr error + + // If no ConnectionIDGenerator is set, this is the ConnectionIDLength. + connIDLen int + // Set in init. + // If no ConnectionIDGenerator is set, this is set to a default. + connIDGenerator ConnectionIDGenerator + statelessResetter *statelessResetter + + server *baseServer + + conn rawConn + + closeQueue chan closePacket + statelessResetQueue chan receivedPacket + + listening chan struct{} // is closed when listen returns + closeErr error + createdConn bool + isSingleUse bool // was created for a single server or client, i.e. by calling quic.Listen or quic.Dial + + readingNonQUICPackets atomic.Bool + nonQUICPackets chan receivedPacket + + logger utils.Logger +} + +// Listen starts listening for incoming QUIC connections. +// There can only be a single listener on any net.PacketConn. +// Listen may only be called again after the current listener was closed. +func (t *Transport) Listen(tlsConf *tls.Config, conf *Config) (*Listener, error) { + s, err := t.createServer(tlsConf, conf, false) + if err != nil { + return nil, err + } + return &Listener{baseServer: s}, nil +} + +// ListenEarly starts listening for incoming QUIC connections. +// There can only be a single listener on any net.PacketConn. +// ListenEarly may only be called again after the current listener was closed. +func (t *Transport) ListenEarly(tlsConf *tls.Config, conf *Config) (*EarlyListener, error) { + s, err := t.createServer(tlsConf, conf, true) + if err != nil { + return nil, err + } + return &EarlyListener{baseServer: s}, nil +} + +func (t *Transport) createServer(tlsConf *tls.Config, conf *Config, allow0RTT bool) (*baseServer, error) { + if tlsConf == nil { + return nil, errors.New("quic: tls.Config not set") + } + if err := validateConfig(conf); err != nil { + return nil, err + } + + t.mutex.Lock() + defer t.mutex.Unlock() + + if t.closeErr != nil { + return nil, t.closeErr + } + if t.server != nil { + return nil, errListenerAlreadySet + } + conf = populateConfig(conf) + if err := t.init(false); err != nil { + return nil, err + } + maxTokenAge := t.MaxTokenAge + if maxTokenAge == 0 { + maxTokenAge = 24 * time.Hour + } + s := newServer( + t.conn, + (*packetHandlerMap)(t), + t.connIDGenerator, + t.statelessResetter, + t.ConnContext, + tlsConf, + conf, + t.Tracer, + t.closeServer, + *t.TokenGeneratorKey, + maxTokenAge, + t.VerifySourceAddress, + t.DisableVersionNegotiationPackets, + allow0RTT, + ) + t.server = s + return s, nil +} + +// Dial dials a new connection to a remote host (not using 0-RTT). +func (t *Transport) Dial(ctx context.Context, addr net.Addr, tlsConf *tls.Config, conf *Config) (*Conn, error) { + return t.dial(ctx, addr, "", tlsConf, conf, false) +} + +// DialEarly dials a new connection, attempting to use 0-RTT if possible. +func (t *Transport) DialEarly(ctx context.Context, addr net.Addr, tlsConf *tls.Config, conf *Config) (*Conn, error) { + return t.dial(ctx, addr, "", tlsConf, conf, true) +} + +func (t *Transport) dial(ctx context.Context, addr net.Addr, host string, tlsConf *tls.Config, conf *Config, use0RTT bool) (*Conn, error) { + if err := t.init(t.isSingleUse); err != nil { + return nil, err + } + if err := validateConfig(conf); err != nil { + return nil, err + } + conf = populateConfig(conf) + tlsConf = tlsConf.Clone() + // setTLSConfigServerName(tlsConf, addr, host) + // The first Initial packet is numbered 1, not 0. + var initialPacketNumber protocol.PacketNumber + if conf.ChromeParrot { + initialPacketNumber = 1 + } + return t.doDial(ctx, + newSendConn(t.conn, addr, packetInfo{}, utils.DefaultLogger), + tlsConf, + conf, + initialPacketNumber, + false, + use0RTT, + conf.Versions[0], + ) +} + +func (t *Transport) doDial( + ctx context.Context, + sendConn sendConn, + tlsConf *tls.Config, + config *Config, + initialPacketNumber protocol.PacketNumber, + hasNegotiatedVersion bool, + use0RTT bool, + version protocol.Version, +) (*Conn, error) { + srcConnID, err := t.connIDGenerator.GenerateConnectionID() + if err != nil { + return nil, err + } + // quic-go randomizes the initial destination connection ID length to exercise + // servers; a fixed length is needed here instead. + genInitialConnID := generateConnectionIDForInitial + if config != nil && config.ChromeParrot { + genInitialConnID = protocol.GenerateChromeConnectionIDForInitial + } + destConnID, err := genInitialConnID() + if err != nil { + return nil, err + } + + t.mutex.Lock() + if t.closeErr != nil { + t.mutex.Unlock() + return nil, t.closeErr + } + + var qlogTrace qlogwriter.Trace + if config.Tracer != nil { + qlogTrace = config.Tracer(ctx, true, destConnID) + } + + logger := utils.DefaultLogger.WithPrefix("client") + logger.Infof("Starting new connection to %s (%s -> %s), source connection ID %s, destination connection ID %s, version %s", tlsConf.ServerName, sendConn.LocalAddr(), sendConn.RemoteAddr(), srcConnID, destConnID, version) + + conn, err := newClientConnection( + context.WithoutCancel(ctx), + sendConn, + (*packetHandlerMap)(t), + destConnID, + srcConnID, + t.connIDGenerator, + t.statelessResetter, + config, + tlsConf, + initialPacketNumber, + use0RTT, + hasNegotiatedVersion, + qlogTrace, + logger, + version, + ) + if err != nil { + t.mutex.Unlock() + return nil, err + } + t.handlers[srcConnID] = conn + t.mutex.Unlock() + + // The error channel needs to be buffered, as the run loop will continue running + // after doDial returns (if the handshake is successful). + // Similarly, the recreateChan needs to be buffered; in case a different case is selected. + errChan := make(chan error, 1) + recreateChan := make(chan errCloseForRecreating, 1) + go func() { + err := conn.run() + var recreateErr *errCloseForRecreating + if errors.As(err, &recreateErr) { + recreateChan <- *recreateErr + return + } + if t.isSingleUse { + t.Close() + } + errChan <- err + }() + + // Only set when we're using 0-RTT. + // Otherwise, earlyConnChan will be nil. Receiving from a nil chan blocks forever. + var earlyConnChan <-chan struct{} + if use0RTT { + earlyConnChan = conn.earlyConnReady() + } + + select { + case <-ctx.Done(): + conn.destroy(nil) + // wait until the Go routine that called Conn.run() returns + select { + case <-errChan: + case <-recreateChan: + } + return nil, context.Cause(ctx) + case params := <-recreateChan: + return t.doDial(ctx, + sendConn, + tlsConf, + config, + params.nextPacketNumber, + true, + use0RTT, + params.nextVersion, + ) + case err := <-errChan: + return nil, err + case <-earlyConnChan: + // ready to send 0-RTT data + return conn.Conn, nil + case <-conn.HandshakeComplete(): + // handshake successfully completed + return conn.Conn, nil + } +} + +func (t *Transport) init(allowZeroLengthConnIDs bool) error { + t.initOnce.Do(func() { + var conn rawConn + if c, ok := t.Conn.(rawConn); ok { + conn = c + } else { + var err error + conn, err = wrapConn(t.Conn, t.DisableGSO) + if err != nil { + t.initErr = err + return + } + } + + t.logger = utils.DefaultLogger // TODO: make this configurable + t.conn = conn + t.handlers = make(map[protocol.ConnectionID]packetHandler) + t.resetTokens = make(map[protocol.StatelessResetToken]packetHandler) + t.listening = make(chan struct{}) + + t.closeQueue = make(chan closePacket, 4) + t.statelessResetQueue = make(chan receivedPacket, 4) + if t.TokenGeneratorKey == nil { + var key TokenGeneratorKey + if _, err := rand.Read(key[:]); err != nil { + t.initErr = err + return + } + t.TokenGeneratorKey = &key + } + + if t.ConnectionIDGenerator != nil { + t.connIDGenerator = t.ConnectionIDGenerator + t.connIDLen = t.ConnectionIDGenerator.ConnectionIDLen() + } else { + connIDLen := t.ConnectionIDLength + if t.ConnectionIDLength == 0 && !allowZeroLengthConnIDs { + connIDLen = protocol.DefaultConnectionIDLength + } + t.connIDLen = connIDLen + t.connIDGenerator = &protocol.DefaultConnectionIDGenerator{ConnLen: t.connIDLen} + } + t.statelessResetter = newStatelessResetter(t.StatelessResetKey) + + go func() { + defer close(t.listening) + t.listen(conn) + + if t.createdConn { + conn.Close() + } + }() + go t.runSendQueue() + }) + return t.initErr +} + +// WriteTo sends a packet on the underlying connection. +func (t *Transport) WriteTo(b []byte, addr net.Addr) (int, error) { + if err := t.init(false); err != nil { + return 0, err + } + return t.conn.WritePacket(b, addr, nil, 0, protocol.ECNUnsupported) +} + +func (t *Transport) runSendQueue() { + for { + select { + case <-t.listening: + return + case p := <-t.closeQueue: + t.conn.WritePacket(p.payload, p.addr, p.info.OOB(), 0, protocol.ECNUnsupported) + case p := <-t.statelessResetQueue: + t.sendStatelessReset(p) + } + } +} + +// Close stops listening for UDP datagrams on the Transport.Conn. +// It abruptly terminates all existing connections, without sending a CONNECTION_CLOSE +// to the peers. It is the application's responsibility to cleanly terminate existing +// connections prior to calling Close. +// +// If a server was started, it will be closed as well. +// It is not possible to start any new server or dial new connections after that. +func (t *Transport) Close() error { + // avoid race condition if the transport is currently being initialized + t.init(false) + + t.close(nil) + if t.createdConn { + if err := t.Conn.Close(); err != nil { + return err + } + } else if t.conn != nil { + t.conn.SetReadDeadline(time.Now()) + defer func() { t.conn.SetReadDeadline(time.Time{}) }() + } + if t.listening != nil { + <-t.listening // wait until listening returns + } + return nil +} + +func (t *Transport) closeServer() { + t.mutex.Lock() + defer t.mutex.Unlock() + + t.server = nil + if t.isSingleUse { + t.closeErr = ErrServerClosed + } + + if len(t.handlers) == 0 { + t.maybeStopListening() + } +} + +func (t *Transport) close(e error) { + t.mutex.Lock() + + if t.closeErr != nil { + t.mutex.Unlock() + return + } + + e = &errTransportClosed{err: e} + t.closeErr = e + server := t.server + t.server = nil + if server != nil { + t.mutex.Unlock() + server.close(e, true) + t.mutex.Lock() + } + + // Close existing connections + var wg sync.WaitGroup + for _, handler := range t.handlers { + wg.Go(func() { handler.destroy(e) }) + } + t.mutex.Unlock() // closing connections requires releasing transport mutex + wg.Wait() + + if t.Tracer != nil { + t.Tracer.Close() + } +} + +func (t *Transport) listen(conn rawConn) { + for { + p, err := conn.ReadPacket() + //nolint:staticcheck // SA1019 ignore this! + // TODO: This code is used to ignore wsa errors on Windows. + // Since net.Error.Temporary is deprecated as of Go 1.18, we should find a better solution. + // See https://github.com/apernet/quic-go/issues/1737 for details. + if nerr, ok := err.(net.Error); ok && nerr.Temporary() { + t.mutex.Lock() + closed := t.closeErr != nil + t.mutex.Unlock() + if closed { + return + } + t.logger.Debugf("Temporary error reading from conn: %w", err) + continue + } + if err != nil { + // Windows returns an error when receiving a UDP datagram that doesn't fit into the provided buffer. + if isRecvMsgSizeErr(err) { + continue + } + t.close(err) + return + } + t.handlePacket(p) + } +} + +func (t *Transport) maybeStopListening() { + if t.isSingleUse && t.closeErr != nil { + t.conn.SetReadDeadline(time.Now()) + } +} + +func (t *Transport) handlePacket(p receivedPacket) { + if len(p.data) == 0 { + return + } + if !wire.IsPotentialQUICPacket(p.data[0]) && !wire.IsLongHeaderPacket(p.data[0]) { + t.handleNonQUICPacket(p) + return + } + connID, err := wire.ParseConnectionID(p.data, t.connIDLen) + if err != nil { + t.logger.Debugf("error parsing connection ID on packet from %s: %s", p.remoteAddr, err) + if t.Tracer != nil { + t.Tracer.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropHeaderParseError, + }) + } + p.buffer.MaybeRelease() + return + } + + // If there's a connection associated with the connection ID, pass the packet there. + if handler, ok := (*packetHandlerMap)(t).Get(connID); ok { + handler.handlePacket(p) + return + } + // RFC 9000 section 10.3.1 requires that the stateless reset detection logic is run for both + // packets that cannot be associated with any connections, and for packets that can't be decrypted. + // We deviate from the RFC and ignore the latter: If a packet's connection ID is associated with an + // existing connection, it is dropped there if if it can't be decrypted. + // Stateless resets use random connection IDs, and at reasonable connection ID lengths collisions are + // exceedingly rare. In the unlikely event that a stateless reset is misrouted to an existing connection, + // it is to be expected that the next stateless reset will be correctly detected. + if isStatelessReset := t.maybeHandleStatelessReset(p.data); isStatelessReset { + return + } + if !wire.IsLongHeaderPacket(p.data[0]) { + if statelessResetQueued := t.maybeSendStatelessReset(p); !statelessResetQueued { + if t.Tracer != nil { + t.Tracer.RecordEvent(qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketType1RTT}, + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnknownConnectionID, + }) + } + p.buffer.Release() + } + return + } + + t.mutex.Lock() + defer t.mutex.Unlock() + if t.server == nil { // no server set + t.logger.Debugf("received a packet with an unexpected connection ID %s", connID) + if t.Tracer != nil { + t.Tracer.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropUnknownConnectionID, + }) + } + p.buffer.MaybeRelease() + return + } + t.server.handlePacket(p) +} + +func (t *Transport) maybeSendStatelessReset(p receivedPacket) (statelessResetQueued bool) { + if t.StatelessResetKey == nil { + return false + } + + // Don't send a stateless reset in response to very small packets. + // This includes packets that could be stateless resets. + if len(p.data) <= protocol.MinStatelessResetSize { + return false + } + + select { + case t.statelessResetQueue <- p: + return true + default: + // it's fine to not send a stateless reset when we're busy + return false + } +} + +func (t *Transport) sendStatelessReset(p receivedPacket) { + defer p.buffer.Release() + + connID, err := wire.ParseConnectionID(p.data, t.connIDLen) + if err != nil { + t.logger.Errorf("error parsing connection ID on packet from %s: %s", p.remoteAddr, err) + return + } + token := t.statelessResetter.GetStatelessResetToken(connID) + t.logger.Debugf("Sending stateless reset to %s (connection ID: %s). Token: %#x", p.remoteAddr, connID, token) + data := make([]byte, protocol.MinStatelessResetSize-16, protocol.MinStatelessResetSize) + rand.Read(data) + data[0] = (data[0] & 0x7f) | 0x40 + data = append(data, token[:]...) + if _, err := t.conn.WritePacket(data, p.remoteAddr, p.info.OOB(), 0, protocol.ECNUnsupported); err != nil { + t.logger.Debugf("Error sending Stateless Reset to %s: %s", p.remoteAddr, err) + } +} + +func (t *Transport) maybeHandleStatelessReset(data []byte) bool { + // stateless resets are always short header packets + if wire.IsLongHeaderPacket(data[0]) { + return false + } + if len(data) < 17 /* type byte + 16 bytes for the reset token */ { + return false + } + + token := protocol.StatelessResetToken(data[len(data)-16:]) + t.mutex.Lock() + conn, ok := t.resetTokens[token] + t.mutex.Unlock() + + if ok { + t.logger.Debugf("Received a stateless reset with token %#x. Closing connection.", token) + go conn.destroy(&StatelessResetError{}) + return true + } + return false +} + +func (t *Transport) handleNonQUICPacket(p receivedPacket) { + // Strictly speaking, this is racy, + // but we only care about receiving packets at some point after ReadNonQUICPacket has been called. + if !t.readingNonQUICPackets.Load() { + return + } + select { + case t.nonQUICPackets <- p: + default: + if t.Tracer != nil { + t.Tracer.RecordEvent(qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: int(p.Size())}, + Trigger: qlog.PacketDropDOSPrevention, + }) + } + } +} + +const maxQueuedNonQUICPackets = 32 + +// ReadNonQUICPacket reads non-QUIC packets received on the underlying connection. +// The detection logic is very simple: Any packet that has the first and second bit of the packet set to 0. +// Note that this is stricter than the detection logic defined in RFC 9443. +func (t *Transport) ReadNonQUICPacket(ctx context.Context, b []byte) (int, net.Addr, error) { + if err := t.init(false); err != nil { + return 0, nil, err + } + if !t.readingNonQUICPackets.Load() { + t.nonQUICPackets = make(chan receivedPacket, maxQueuedNonQUICPackets) + t.readingNonQUICPackets.Store(true) + } + select { + case <-ctx.Done(): + return 0, nil, ctx.Err() + case p := <-t.nonQUICPackets: + n := copy(b, p.data) + return n, p.remoteAddr, nil + case <-t.listening: + return 0, nil, errors.New("closed") + } +} + +func setTLSConfigServerName(tlsConf *tls.Config, addr net.Addr, host string) { + // If no ServerName is set, infer the ServerName from the host we're connecting to. + if tlsConf.ServerName != "" { + return + } + if host == "" { + if udpAddr, ok := addr.(*net.UDPAddr); ok { + tlsConf.ServerName = udpAddr.IP.String() + return + } + } + h, _, err := net.SplitHostPort(host) + if err != nil { // This happens if the host doesn't contain a port number. + tlsConf.ServerName = host + return + } + tlsConf.ServerName = h +} + +type packetHandlerMap Transport + +var _ connRunner = &packetHandlerMap{} + +func (h *packetHandlerMap) Add(id protocol.ConnectionID, handler packetHandler) bool /* was added */ { + h.mutex.Lock() + defer h.mutex.Unlock() + + if _, ok := h.handlers[id]; ok { + h.logger.Debugf("Not adding connection ID %s, as it already exists.", id) + return false + } + h.handlers[id] = handler + h.logger.Debugf("Adding connection ID %s.", id) + return true +} + +func (h *packetHandlerMap) Get(connID protocol.ConnectionID) (packetHandler, bool) { + h.mutex.Lock() + defer h.mutex.Unlock() + handler, ok := h.handlers[connID] + return handler, ok +} + +func (h *packetHandlerMap) AddResetToken(token protocol.StatelessResetToken, handler packetHandler) { + h.mutex.Lock() + h.resetTokens[token] = handler + h.mutex.Unlock() +} + +func (h *packetHandlerMap) RemoveResetToken(token protocol.StatelessResetToken) { + h.mutex.Lock() + delete(h.resetTokens, token) + h.mutex.Unlock() +} + +func (h *packetHandlerMap) AddWithConnID(clientDestConnID, newConnID protocol.ConnectionID, handler packetHandler) bool { + h.mutex.Lock() + defer h.mutex.Unlock() + + if _, ok := h.handlers[clientDestConnID]; ok { + h.logger.Debugf("Not adding connection ID %s for a new connection, as it already exists.", clientDestConnID) + return false + } + h.handlers[clientDestConnID] = handler + h.handlers[newConnID] = handler + h.logger.Debugf("Adding connection IDs %s and %s for a new connection.", clientDestConnID, newConnID) + return true +} + +func (h *packetHandlerMap) Remove(id protocol.ConnectionID) { + h.mutex.Lock() + delete(h.handlers, id) + h.mutex.Unlock() + h.logger.Debugf("Removing connection ID %s.", id) +} + +// ReplaceWithClosed is called when a connection is closed. +// Depending on which side closed the connection, we need to: +// * remote close: absorb delayed packets +// * local close: retransmit the CONNECTION_CLOSE packet, in case it was lost +func (h *packetHandlerMap) ReplaceWithClosed(ids []protocol.ConnectionID, connClosePacket []byte, expiry time.Duration) { + var handler packetHandler + if connClosePacket != nil { + handler = newClosedLocalConn( + func(addr net.Addr, info packetInfo) { + select { + case h.closeQueue <- closePacket{payload: connClosePacket, addr: addr, info: info}: + default: + // We're backlogged. + // Just drop the packet, sending CONNECTION_CLOSE copies is best effort anyway. + } + }, + h.logger, + ) + } else { + handler = newClosedRemoteConn() + } + + h.mutex.Lock() + for _, id := range ids { + h.handlers[id] = handler + } + h.mutex.Unlock() + h.logger.Debugf("Replacing connection for connection IDs %s with a closed connection.", ids) + + time.AfterFunc(expiry, func() { + h.mutex.Lock() + for _, id := range ids { + delete(h.handlers, id) + } + if len(h.handlers) == 0 { + t := (*Transport)(h) + t.maybeStopListening() + } + h.mutex.Unlock() + h.logger.Debugf("Removing connection IDs %s for a closed connection after it has been retired.", ids) + }) +} diff --git a/third_party/quic-go/transport_test.go b/third_party/quic-go/transport_test.go new file mode 100644 index 0000000..65d0613 --- /dev/null +++ b/third_party/quic-go/transport_test.go @@ -0,0 +1,739 @@ +package quic + +import ( + "bytes" + "context" + "crypto/tls" + "errors" + "math" + "net" + "sync/atomic" + "syscall" + "testing" + "testing/synctest" + "time" + + "github.com/apernet/quic-go/internal/protocol" + "github.com/apernet/quic-go/internal/qerr" + "github.com/apernet/quic-go/internal/utils" + "github.com/apernet/quic-go/internal/wire" + "github.com/apernet/quic-go/qlog" + "github.com/apernet/quic-go/qlogwriter" + "github.com/apernet/quic-go/testutils/events" + "github.com/apernet/quic-go/testutils/simnet" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type mockPacketConn struct { + localAddr net.Addr + readErrs chan error +} + +func (c *mockPacketConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + err, ok := <-c.readErrs + if !ok { + return 0, nil, net.ErrClosed + } + return 0, nil, err +} + +func (c *mockPacketConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { panic("implement me") } +func (c *mockPacketConn) LocalAddr() net.Addr { return c.localAddr } +func (c *mockPacketConn) Close() error { close(c.readErrs); return nil } +func (c *mockPacketConn) SetDeadline(t time.Time) error { return nil } +func (c *mockPacketConn) SetReadDeadline(t time.Time) error { return nil } +func (c *mockPacketConn) SetWriteDeadline(t time.Time) error { return nil } + +type mockPacketHandler struct { + packets chan<- receivedPacket + destruction chan<- error +} + +func (h *mockPacketHandler) handlePacket(p receivedPacket) { + h.packets <- p +} + +func (h *mockPacketHandler) destroy(err error) { + if h.destruction != nil { + h.destruction <- err + } +} + +func (h *mockPacketHandler) closeWithTransportError(code qerr.TransportErrorCode) {} + +func newSimnetLink(t *testing.T, rtt time.Duration) (client, server net.PacketConn, close func()) { + t.Helper() + + n := &simnet.Simnet{Router: &simnet.PerfectRouter{}} + settings := simnet.NodeBiDiLinkSettings{Latency: rtt / 2} + + client = n.NewEndpoint(&net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 9001}, settings) + server = n.NewEndpoint(&net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 9002}, settings) + require.NoError(t, n.Start()) + return client, server, func() { + require.NoError(t, n.Close()) + } +} + +func TestTransportPacketHandling(t *testing.T) { + tr := &Transport{Conn: newUDPConnLocalhost(t)} + tr.init(true) + defer tr.Close() + + connID1 := protocol.ParseConnectionID([]byte{1, 2, 3, 4, 5, 6, 7, 8}) + connID2 := protocol.ParseConnectionID([]byte{8, 7, 6, 5, 4, 3, 2, 1}) + + connChan1 := make(chan receivedPacket, 1) + conn1 := &mockPacketHandler{packets: connChan1} + (*packetHandlerMap)(tr).Add(connID1, conn1) + connChan2 := make(chan receivedPacket, 1) + conn2 := &mockPacketHandler{packets: connChan2} + (*packetHandlerMap)(tr).Add(connID2, conn2) + + conn := newUDPConnLocalhost(t) + _, err := conn.WriteTo(getPacket(t, connID1), tr.Conn.LocalAddr()) + require.NoError(t, err) + _, err = conn.WriteTo(getPacket(t, connID2), tr.Conn.LocalAddr()) + require.NoError(t, err) + + select { + case p := <-connChan1: + require.Equal(t, conn.LocalAddr(), p.remoteAddr) + connID, err := wire.ParseConnectionID(p.data, 0) + require.NoError(t, err) + require.Equal(t, connID1, connID) + case <-time.After(time.Second): + t.Fatal("timeout") + } + select { + case p := <-connChan2: + require.Equal(t, conn.LocalAddr(), p.remoteAddr) + connID, err := wire.ParseConnectionID(p.data, 0) + require.NoError(t, err) + require.Equal(t, connID2, connID) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportAndListenerConcurrentClose(t *testing.T) { + tr := &Transport{Conn: newUDPConnLocalhost(t)} + ln, err := tr.Listen(&tls.Config{}, nil) + require.NoError(t, err) + // close transport and listener concurrently + lnErrChan := make(chan error, 1) + go func() { lnErrChan <- ln.Close() }() + require.NoError(t, tr.Close()) + select { + case err := <-lnErrChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportAndDialConcurrentClose(t *testing.T) { + server := newUDPConnLocalhost(t) + + tr := &Transport{Conn: newUDPConnLocalhost(t)} + // close transport and dial concurrently + errChan := make(chan error, 1) + go func() { errChan <- tr.Close() }() + ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) + defer cancel() + _, err := tr.Dial(ctx, server.LocalAddr(), &tls.Config{InsecureSkipVerify: true}, nil) + require.Error(t, err) + require.ErrorIs(t, err, ErrTransportClosed) + require.NotErrorIs(t, err, context.DeadlineExceeded) + + select { + case <-errChan: + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportErrFromConn(t *testing.T) { + t.Setenv("QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING", "true") + + synctest.Test(t, func(t *testing.T) { + readErrChan := make(chan error, 2) + tr := Transport{ + Conn: &mockPacketConn{ + readErrs: readErrChan, + localAddr: &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 1234}, + }, + } + defer tr.Close() + tr.init(true) + + errChan := make(chan error, 1) + ph := &mockPacketHandler{destruction: errChan} + (*packetHandlerMap)(&tr).Add(protocol.ParseConnectionID([]byte{1, 2, 3, 4}), ph) + + // temporary errors don't lead to a shutdown... + var tempErr deadlineError + require.True(t, tempErr.Temporary()) + readErrChan <- tempErr + // don't expect any calls to phm.Close + synctest.Wait() + + // ...but non-temporary errors do + readErrChan <- errors.New("read failed") + synctest.Wait() + + select { + case err := <-errChan: + require.ErrorIs(t, err, ErrTransportClosed) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + _, err := tr.Listen(&tls.Config{}, nil) + require.ErrorIs(t, err, ErrTransportClosed) + }) +} + +func TestTransportStatelessResetReceiving(t *testing.T) { + tr := &Transport{ + Conn: newUDPConnLocalhost(t), + ConnectionIDLength: 4, + } + tr.init(true) + defer tr.Close() + + connID := protocol.ParseConnectionID([]byte{9, 10, 11, 12}) + // now send a packet with a connection ID that doesn't exist + token := protocol.StatelessResetToken{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16} + b, err := wire.AppendShortHeader(nil, connID, 1337, 2, protocol.KeyPhaseOne) + require.NoError(t, err) + b = append(b, token[:]...) + + destroyChan := make(chan error, 1) + conn1 := &mockPacketHandler{destruction: destroyChan} + (*packetHandlerMap)(tr).AddResetToken(token, conn1) + + conn := newUDPConnLocalhost(t) + _, err = conn.WriteTo(b, tr.Conn.LocalAddr()) + require.NoError(t, err) + + select { + case err := <-destroyChan: + require.ErrorIs(t, err, &qerr.StatelessResetError{}) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportStatelessResetSending(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + clientConn, serverConn, closeFn := newSimnetLink(t, rtt) + defer closeFn() + + var eventRecorder events.Recorder + tr := &Transport{ + Conn: serverConn, + ConnectionIDLength: 4, + StatelessResetKey: &StatelessResetKey{1, 2, 3, 4}, + Tracer: &eventRecorder, + } + tr.init(true) + defer tr.Close() + + connID := protocol.ParseConnectionID([]byte{9, 10, 11, 12}) + + // now send a packet with a connection ID that doesn't exist + b, err := wire.AppendShortHeader(nil, connID, 1337, 2, protocol.KeyPhaseOne) + require.NoError(t, err) + + // no stateless reset sent for packets smaller than MinStatelessResetSize + smallPacket := append(b, make([]byte, protocol.MinStatelessResetSize-len(b))...) + _, err = clientConn.WriteTo(smallPacket, tr.Conn.LocalAddr()) + require.NoError(t, err) + + time.Sleep(rtt) // so that the packet arrives at the server + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketType1RTT}, + Raw: qlog.RawInfo{Length: len(smallPacket)}, + Trigger: qlog.PacketDropUnknownConnectionID, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + + // but a stateless reset is sent for packets larger than MinStatelessResetSize + _, err = clientConn.WriteTo(append(b, make([]byte, protocol.MinStatelessResetSize-len(b)+1)...), tr.Conn.LocalAddr()) + require.NoError(t, err) + clientConn.SetReadDeadline(time.Now().Add(time.Second)) + p := make([]byte, 1024) + n, addr, err := clientConn.ReadFrom(p) + require.NoError(t, err) + require.Equal(t, addr, tr.Conn.LocalAddr()) + srt := newStatelessResetter(tr.StatelessResetKey).GetStatelessResetToken(connID) + require.Contains(t, string(p[:n]), string(srt[:])) + }) +} + +func TestTransportUnparseableQUICPackets(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + clientConn, serverConn, closeFn := newSimnetLink(t, rtt) + defer closeFn() + + var eventRecorder events.Recorder + tr := &Transport{ + Conn: serverConn, + ConnectionIDLength: 10, + Tracer: &eventRecorder, + } + require.NoError(t, tr.init(true)) + defer tr.Close() + + _, err := clientConn.WriteTo([]byte{0x40 /* set the QUIC bit */, 1, 2, 3}, tr.Conn.LocalAddr()) + require.NoError(t, err) + + time.Sleep(rtt) // so that the packet arrives at the server + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: 4}, + Trigger: qlog.PacketDropHeaderParseError, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + }) +} + +func TestTransportListening(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + clientConn, serverConn, closeFn := newSimnetLink(t, rtt) + defer closeFn() + + var eventRecorder events.Recorder + tr := &Transport{ + Conn: serverConn, + ConnectionIDLength: 5, + Tracer: &eventRecorder, + } + require.NoError(t, tr.init(true)) + defer tr.Close() + + data := wire.ComposeVersionNegotiation([]byte{1, 2, 3, 4, 5}, []byte{6, 7, 8, 9, 10}, []protocol.Version{protocol.Version1}) + + _, err := clientConn.WriteTo(data, tr.Conn.LocalAddr()) + require.NoError(t, err) + + time.Sleep(rtt) // so that the packet arrives at the server + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Raw: qlog.RawInfo{Length: len(data)}, + Trigger: qlog.PacketDropUnknownConnectionID, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + eventRecorder.Clear() + + ln, err := tr.Listen(&tls.Config{}, nil) + require.NoError(t, err) + + _, err = clientConn.WriteTo(data, tr.Conn.LocalAddr()) + require.NoError(t, err) + time.Sleep(rtt) // so that the packet arrives at the server + + require.Equal(t, + []qlogwriter.Event{ + qlog.PacketDropped{ + Header: qlog.PacketHeader{PacketType: qlog.PacketTypeVersionNegotiation}, + Raw: qlog.RawInfo{Length: len(data)}, + Trigger: qlog.PacketDropUnexpectedPacket, + }, + }, + eventRecorder.Events(qlog.PacketDropped{}), + ) + + // only a single listener can be set + _, err = tr.Listen(&tls.Config{}, nil) + require.Error(t, err) + require.ErrorIs(t, err, errListenerAlreadySet) + + require.NoError(t, ln.Close()) + // now it's possible to add a new listener + ln, err = tr.Listen(&tls.Config{}, nil) + require.NoError(t, err) + defer ln.Close() + }) +} + +func TestTransportNonQUICPackets(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const rtt = 10 * time.Millisecond + clientConn, serverConn, closeFn := newSimnetLink(t, rtt) + defer closeFn() + + tr := &Transport{Conn: serverConn} + defer tr.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Millisecond) + defer cancel() + _, _, err := tr.ReadNonQUICPacket(ctx, make([]byte, 1024)) + require.Error(t, err) + require.ErrorIs(t, err, context.DeadlineExceeded) + + data := []byte{0 /* don't set the QUIC bit */, 1, 2, 3} + _, err = clientConn.WriteTo(data, tr.Conn.LocalAddr()) + require.NoError(t, err) + _, err = clientConn.WriteTo(data, tr.Conn.LocalAddr()) + require.NoError(t, err) + + ctx, cancel = context.WithTimeout(context.Background(), time.Second) + defer cancel() + b := make([]byte, 1024) + n, addr, err := tr.ReadNonQUICPacket(ctx, b) + require.NoError(t, err) + require.Equal(t, data, b[:n]) + require.Equal(t, addr, clientConn.LocalAddr()) + + // now send a lot of packets without reading them + for i := range 2 * maxQueuedNonQUICPackets { + data := append([]byte{0 /* don't set the QUIC bit */, uint8(i)}, bytes.Repeat([]byte{uint8(i)}, 1000)...) + _, err = clientConn.WriteTo(data, tr.Conn.LocalAddr()) + require.NoError(t, err) + } + + time.Sleep(rtt) // so that all packets arrive at the server + + var received int + for { + ctx, cancel = context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + _, _, err := tr.ReadNonQUICPacket(ctx, b) + if errors.Is(err, context.DeadlineExceeded) { + break + } + require.NoError(t, err) + received++ + } + require.Equal(t, received, maxQueuedNonQUICPackets) + }) +} + +type faultySyscallConn struct{ net.PacketConn } + +func (c *faultySyscallConn) SyscallConn() (syscall.RawConn, error) { return nil, assert.AnError } + +func TestTransportFaultySyscallConn(t *testing.T) { + syscallconn := &faultySyscallConn{PacketConn: newUDPConnLocalhost(t)} + + tr := &Transport{Conn: syscallconn} + _, err := tr.Listen(&tls.Config{}, nil) + require.Error(t, err) + require.ErrorIs(t, err, assert.AnError) +} + +func TestTransportSetTLSConfigServerName(t *testing.T) { + for _, tt := range []struct { + name string + expected string + conf *tls.Config + host string + }{ + { + name: "uses the value from the config", + expected: "foo.bar", + conf: &tls.Config{ServerName: "foo.bar"}, + host: "baz.foo", + }, + { + name: "uses the hostname", + expected: "golang.org", + conf: &tls.Config{}, + host: "golang.org", + }, + { + name: "removes the port from the hostname", + expected: "golang.org", + conf: &tls.Config{}, + host: "golang.org:1234", + }, + { + name: "uses the IP", + expected: "1.3.5.7", + conf: &tls.Config{}, + host: "", + }, + } { + t.Run(tt.name, func(t *testing.T) { + setTLSConfigServerName(tt.conf, &net.UDPAddr{IP: net.IPv4(1, 3, 5, 7), Port: 1234}, tt.host) + require.Equal(t, tt.expected, tt.conf.ServerName) + }) + } +} + +func TestTransportDial(t *testing.T) { + t.Run("regular", func(t *testing.T) { + testTransportDial(t, false) + }) + + t.Run("early", func(t *testing.T) { + testTransportDial(t, true) + }) +} + +func testTransportDial(t *testing.T, early bool) { + originalClientConnConstructor := newClientConnection + t.Cleanup(func() { newClientConnection = originalClientConnConstructor }) + + synctest.Test(t, func(t *testing.T) { + _, serverConn, closeFn := newSimnetLink(t, 10*time.Millisecond) + defer closeFn() + + var conn *connTestHooks + handshakeChan := make(chan struct{}) + blockRun := make(chan struct{}) + if early { + conn = &connTestHooks{ + earlyConnReady: func() <-chan struct{} { return handshakeChan }, + handshakeComplete: func() <-chan struct{} { return make(chan struct{}) }, + } + } else { + conn = &connTestHooks{ + handshakeComplete: func() <-chan struct{} { return handshakeChan }, + } + } + conn.run = func() error { <-blockRun; return errors.New("done") } + defer close(blockRun) + + newClientConnection = func( + _ context.Context, + _ sendConn, + _ connRunner, + _ protocol.ConnectionID, + _ protocol.ConnectionID, + _ ConnectionIDGenerator, + _ *statelessResetter, + _ *Config, + _ *tls.Config, + _ protocol.PacketNumber, + _ bool, + _ bool, + _ qlogwriter.Trace, + _ utils.Logger, + _ protocol.Version, + ) (*wrappedConn, error) { + return &wrappedConn{testHooks: conn}, nil + } + + tr := &Transport{Conn: serverConn} + tr.init(true) + defer tr.Close() + + errChan := make(chan error, 1) + go func() { + var err error + if early { + _, err = tr.DialEarly(context.Background(), nil, &tls.Config{}, nil) + } else { + _, err = tr.Dial(context.Background(), nil, &tls.Config{}, nil) + } + errChan <- err + }() + + synctest.Wait() + + select { + case <-errChan: + t.Fatal("Dial shouldn't have returned") + default: + } + + close(handshakeChan) + + synctest.Wait() + + select { + case err := <-errChan: + require.NoError(t, err) + default: + } + }) +} + +func TestTransportDialingVersionNegotiation(t *testing.T) { + originalClientConnConstructor := newClientConnection + t.Cleanup(func() { newClientConnection = originalClientConnConstructor }) + + conn := &connTestHooks{ + handshakeComplete: func() <-chan struct{} { return make(chan struct{}) }, + run: func() error { return &errCloseForRecreating{nextPacketNumber: 109, nextVersion: 789} }, + } + conn2 := &connTestHooks{ + handshakeComplete: func() <-chan struct{} { return make(chan struct{}) }, + run: func() error { return assert.AnError }, + } + + type connParams struct { + pn protocol.PacketNumber + hasNegotiatedVersion bool + version protocol.Version + } + + connChan := make(chan connParams, 2) + var counter int + newClientConnection = func( + _ context.Context, + _ sendConn, + _ connRunner, + _ protocol.ConnectionID, + _ protocol.ConnectionID, + _ ConnectionIDGenerator, + _ *statelessResetter, + _ *Config, + _ *tls.Config, + pn protocol.PacketNumber, + _ bool, + hasNegotiatedVersion bool, + _ qlogwriter.Trace, + _ utils.Logger, + v protocol.Version, + ) (*wrappedConn, error) { + connChan <- connParams{pn: pn, hasNegotiatedVersion: hasNegotiatedVersion, version: v} + if counter == 0 { + counter++ + return &wrappedConn{testHooks: conn}, nil + } + return &wrappedConn{testHooks: conn2}, nil + } + + tr := &Transport{Conn: newUDPConnLocalhost(t)} + tr.init(true) + defer tr.Close() + + _, err := tr.Dial(context.Background(), nil, &tls.Config{}, nil) + require.ErrorIs(t, err, assert.AnError) + + select { + case params := <-connChan: + require.Zero(t, params.pn) + require.False(t, params.hasNegotiatedVersion) + require.Equal(t, protocol.Version1, params.version) + case <-time.After(time.Second): + t.Fatal("timeout") + } + select { + case params := <-connChan: + require.Equal(t, protocol.PacketNumber(109), params.pn) + require.True(t, params.hasNegotiatedVersion) + require.Equal(t, protocol.Version(789), params.version) + case <-time.After(time.Second): + t.Fatal("timeout") + } +} + +func TestTransportReplaceWithClosed(t *testing.T) { + t.Run("local", func(t *testing.T) { + testTransportReplaceWithClosed(t, true) + }) + t.Run("remote", func(t *testing.T) { + testTransportReplaceWithClosed(t, false) + }) +} + +func testTransportReplaceWithClosed(t *testing.T, local bool) { + synctest.Test(t, func(t *testing.T) { + clientConn, serverConn, closeFn := newSimnetLink(t, 10*time.Millisecond) + defer closeFn() + + srk := StatelessResetKey{1, 2, 3, 4} + tr := &Transport{ + Conn: serverConn, + ConnectionIDLength: 4, + StatelessResetKey: &srk, + } + tr.init(true) + defer tr.Close() + + var closePacket []byte + if local { + closePacket = []byte("foobar") + } + + const expiry = 50 * time.Millisecond + handler := &mockPacketHandler{} + connID := protocol.ParseConnectionID([]byte{4, 3, 2, 1}) + m := (*packetHandlerMap)(tr) + require.True(t, m.Add(connID, handler)) + m.ReplaceWithClosed([]protocol.ConnectionID{connID}, closePacket, expiry) + + p := make([]byte, 100) + p[0] = 0x40 // QUIC bit + copy(p[1:], connID.Bytes()) + + var sent atomic.Int64 + errChan := make(chan error, 1) + stopSending := make(chan struct{}) + go func() { + defer close(errChan) + ticker := time.NewTicker(expiry / 200) + timeout := time.NewTimer(time.Second) + for { + select { + case <-stopSending: + return + case <-timeout.C: + errChan <- errors.New("timeout") + return + case <-ticker.C: + } + if _, err := clientConn.WriteTo(p, tr.Conn.LocalAddr()); err != nil { + errChan <- err + return + } + sent.Add(1) + } + }() + + // For locally closed connections, CONNECTION_CLOSE packets are sent with an exponential backoff + var received int + clientConn.SetReadDeadline(time.Now().Add(time.Hour)) + for { + b := make([]byte, 100) + n, _, err := clientConn.ReadFrom(b) + require.NoError(t, err) + // at some point, the connection is cleaned up, and we'll receive a stateless reset + if !bytes.Equal(b[:n], []byte("foobar")) { + require.GreaterOrEqual(t, n, protocol.MinStatelessResetSize) + close(stopSending) // stop sending packets + break + } + received++ + } + + select { + case err := <-errChan: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("timeout") + } + + numSent := sent.Load() + if !local { + require.Zero(t, received) + t.Logf("sent %d packets", numSent) + return + } + t.Logf("sent %d packets, received %d CONNECTION_CLOSE copies", numSent, received) + require.Equal(t, int(math.Ceil(math.Log2(float64(numSent)))), received) + }) +} diff --git a/tools/notices/main.go b/tools/notices/main.go new file mode 100644 index 0000000..4a5e837 --- /dev/null +++ b/tools/notices/main.go @@ -0,0 +1,320 @@ +// Command notices generates THIRD_PARTY_NOTICES.md from the modules that are +// actually present in AutoCAR's linked package graph. +package main + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "flag" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "sort" + "strings" + "unicode/utf8" +) + +const defaultOutput = "THIRD_PARTY_NOTICES.md" + +type listedPackage struct { + Module *listedModule `json:"Module"` +} + +type listedModule struct { + Path string `json:"Path"` + Version string `json:"Version"` + Dir string `json:"Dir"` + Main bool `json:"Main"` + Replace *listedModule `json:"Replace"` +} + +type dependency struct { + Path string + Version string + Replacement string + Dir string + Notices []noticeFile +} + +type noticeFile struct { + Name string + Text string + SHA256 string +} + +func main() { + var ( + check = flag.Bool("check", false, "verify that the output file is current") + output = flag.String("out", defaultOutput, "output path, relative to the main module") + target = flag.String("target", "./cmd/autocar", "Go package whose linked dependencies are inspected") + ) + flag.Parse() + if flag.NArg() != 0 { + fmt.Fprintln(os.Stderr, "notices: unexpected positional arguments") + os.Exit(2) + } + + if err := run(context.Background(), *target, *output, *check); err != nil { + fmt.Fprintln(os.Stderr, "notices:", err) + os.Exit(1) + } +} + +func run(ctx context.Context, target, output string, check bool) error { + root, err := moduleRoot(ctx) + if err != nil { + return err + } + dependencies, err := collectDependencies(ctx, root, target) + if err != nil { + return err + } + generated, err := render(dependencies) + if err != nil { + return err + } + + outputPath := output + if !filepath.IsAbs(outputPath) { + outputPath = filepath.Join(root, outputPath) + } + if check { + current, readErr := os.ReadFile(outputPath) + if readErr != nil { + return fmt.Errorf("read %s: %w", filepath.Base(outputPath), readErr) + } + if !bytes.Equal(current, generated) { + return fmt.Errorf("%s is stale; run `go run ./tools/notices`", filepath.Base(outputPath)) + } + fmt.Printf("%s is current (%d linked modules)\n", filepath.Base(outputPath), len(dependencies)) + return nil + } + if err := os.WriteFile(outputPath, generated, 0o644); err != nil { + return fmt.Errorf("write %s: %w", filepath.Base(outputPath), err) + } + fmt.Printf("wrote %s (%d linked modules)\n", filepath.Base(outputPath), len(dependencies)) + return nil +} + +func moduleRoot(ctx context.Context) (string, error) { + command := exec.CommandContext(ctx, "go", "env", "GOMOD") + output, err := command.Output() + if err != nil { + return "", fmt.Errorf("locate main module: %w", err) + } + goMod := strings.TrimSpace(string(output)) + if goMod == "" || goMod == os.DevNull { + return "", errors.New("not running in a Go module") + } + return filepath.Dir(goMod), nil +} + +func collectDependencies(ctx context.Context, root, target string) ([]dependency, error) { + command := exec.CommandContext(ctx, "go", "list", "-deps", "-json", target) + command.Dir = root + output, err := command.Output() + if err != nil { + var exitErr *exec.ExitError + if errors.As(err, &exitErr) { + return nil, fmt.Errorf("go list %s: %s", target, strings.TrimSpace(string(exitErr.Stderr))) + } + return nil, fmt.Errorf("go list %s: %w", target, err) + } + + unique := make(map[string]dependency) + decoder := json.NewDecoder(bytes.NewReader(output)) + for { + var pkg listedPackage + if err := decoder.Decode(&pkg); errors.Is(err, io.EOF) { + break + } else if err != nil { + return nil, fmt.Errorf("decode go list output: %w", err) + } + if pkg.Module == nil || pkg.Module.Main { + continue + } + + item, err := dependencyFromModule(root, pkg.Module) + if err != nil { + return nil, err + } + key := item.Path + "\x00" + item.Version + "\x00" + item.Replacement + if previous, ok := unique[key]; ok { + if previous.Dir != item.Dir { + return nil, fmt.Errorf("module %s resolved to multiple directories", item.Path) + } + continue + } + unique[key] = item + } + + dependencies := make([]dependency, 0, len(unique)) + for _, item := range unique { + notices, err := readNotices(item.Dir) + if err != nil { + return nil, fmt.Errorf("module %s: %w", item.Path, err) + } + item.Notices = notices + dependencies = append(dependencies, item) + } + sort.Slice(dependencies, func(i, j int) bool { + if dependencies[i].Path != dependencies[j].Path { + return dependencies[i].Path < dependencies[j].Path + } + if dependencies[i].Version != dependencies[j].Version { + return dependencies[i].Version < dependencies[j].Version + } + return dependencies[i].Replacement < dependencies[j].Replacement + }) + if len(dependencies) == 0 { + return nil, errors.New("no non-main modules found in linked package graph") + } + return dependencies, nil +} + +func dependencyFromModule(root string, module *listedModule) (dependency, error) { + item := dependency{Path: module.Path, Version: module.Version, Dir: module.Dir} + if module.Replace != nil { + item.Dir = module.Replace.Dir + item.Replacement = replacementLabel(root, module.Replace) + } + if item.Path == "" { + return dependency{}, errors.New("go list returned a module without a path") + } + if item.Dir == "" { + return dependency{}, fmt.Errorf("module %s has no resolved source directory", item.Path) + } + return item, nil +} + +func replacementLabel(root string, replacement *listedModule) string { + if replacement == nil { + return "" + } + path := filepath.ToSlash(replacement.Path) + if filepath.IsAbs(replacement.Path) { + if relative, err := filepath.Rel(root, replacement.Path); err == nil && relative != ".." && !strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + path = "./" + filepath.ToSlash(relative) + } else { + path = "local replacement" + } + } + if replacement.Version != "" { + return strings.TrimSpace(path + " " + replacement.Version) + } + return path +} + +func readNotices(dir string) ([]noticeFile, error) { + entries, err := os.ReadDir(dir) + if err != nil { + return nil, fmt.Errorf("read module root: %w", err) + } + var names []string + for _, entry := range entries { + if !isNoticeName(entry.Name()) { + continue + } + info, err := os.Stat(filepath.Join(dir, entry.Name())) + if err != nil { + return nil, fmt.Errorf("stat %s: %w", entry.Name(), err) + } + if info.Mode().IsRegular() { + names = append(names, entry.Name()) + } + } + sort.Strings(names) + if len(names) == 0 { + return nil, errors.New("no root LICENSE, COPYING, NOTICE, or PATENTS file found") + } + + notices := make([]noticeFile, 0, len(names)) + for _, name := range names { + raw, err := os.ReadFile(filepath.Join(dir, name)) + if err != nil { + return nil, fmt.Errorf("read %s: %w", name, err) + } + if !utf8.Valid(raw) || bytes.IndexByte(raw, 0) >= 0 { + return nil, fmt.Errorf("%s is not a UTF-8 text file", name) + } + digest := sha256.Sum256(raw) + notices = append(notices, noticeFile{ + Name: name, + Text: normalizeNewlines(string(raw)), + SHA256: hex.EncodeToString(digest[:]), + }) + } + return notices, nil +} + +func isNoticeName(name string) bool { + lower := strings.ToLower(name) + for _, prefix := range []string{"license", "copying", "notice", "patents"} { + if lower == prefix || strings.HasPrefix(lower, prefix+".") || strings.HasPrefix(lower, prefix+"-") || strings.HasPrefix(lower, prefix+"_") { + return true + } + } + return false +} + +func normalizeNewlines(value string) string { + value = strings.ReplaceAll(value, "\r\n", "\n") + value = strings.ReplaceAll(value, "\r", "\n") + return strings.TrimRight(value, "\n") + "\n" +} + +func render(dependencies []dependency) ([]byte, error) { + var output bytes.Buffer + output.WriteString("# Third-party notices\n\n") + output.WriteString("> Generated by `go run ./tools/notices`; do not edit by hand. CI verifies this file against the linked build graph.\n\n") + output.WriteString("AutoCAR itself is licensed under the repository's `LICENSE`. The sections below reproduce every root-level license, notice, and patent file from each non-main Go module reached by `go list -deps -json ./cmd/autocar`. A module-level replacement is recorded so the notice always describes the source that is actually compiled.\n\n") + output.WriteString("The SHA-256 value is calculated from the upstream file's original bytes; line endings in the displayed copy are normalized for Markdown.\n") + + for _, item := range dependencies { + if item.Path == "" || len(item.Notices) == 0 { + return nil, errors.New("cannot render incomplete dependency metadata") + } + fmt.Fprintf(&output, "\n## `%s`", item.Path) + if item.Version != "" { + fmt.Fprintf(&output, " `%s`", item.Version) + } + output.WriteString("\n\n") + if item.Replacement != "" { + fmt.Fprintf(&output, "Effective source replacement: `%s`.\n\n", item.Replacement) + } + for _, notice := range item.Notices { + fmt.Fprintf(&output, "### `%s`\n\nSHA-256: `%s`\n\n", notice.Name, notice.SHA256) + fence := markdownFence(notice.Text) + fmt.Fprintf(&output, "%stext\n%s%s\n", fence, notice.Text, fence) + } + } + + output.WriteString("\n## Brutal code provenance\n\n") + output.WriteString("AutoCAR's negotiated Brutal controller is provided by the MIT-licensed Hysteria core module identified above. AutoCAR does not vendor, import, or copy the GPL-licensed `tcp-brutal` implementation. Similar terminology describes a traffic-control strategy and does not imply source-code provenance.\n") + return output.Bytes(), nil +} + +func markdownFence(value string) string { + longest, current := 0, 0 + for _, character := range value { + if character == '`' { + current++ + if current > longest { + longest = current + } + } else { + current = 0 + } + } + length := 3 + if longest >= length { + length = longest + 1 + } + return strings.Repeat("`", length) +} diff --git a/tools/notices/main_test.go b/tools/notices/main_test.go new file mode 100644 index 0000000..7b2f69e --- /dev/null +++ b/tools/notices/main_test.go @@ -0,0 +1,87 @@ +package main + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestReadNoticesFindsAndSortsRootFiles(t *testing.T) { + directory := t.TempDir() + for name, content := range map[string]string{ + "NOTICE": "notice\r\n", + "LICENSE-MIT": "license\n", + "PATENTS.txt": "patents\n", + "README.md": "not a notice\n", + "UNLICENSED.md": "not a notice\n", + } { + if err := os.WriteFile(filepath.Join(directory, name), []byte(content), 0o644); err != nil { + t.Fatal(err) + } + } + + notices, err := readNotices(directory) + if err != nil { + t.Fatal(err) + } + want := []string{"LICENSE-MIT", "NOTICE", "PATENTS.txt"} + if len(notices) != len(want) { + t.Fatalf("notice count = %d, want %d", len(notices), len(want)) + } + for index := range want { + if notices[index].Name != want[index] { + t.Fatalf("notice[%d] = %q, want %q", index, notices[index].Name, want[index]) + } + if strings.Contains(notices[index].Text, "\r") || !strings.HasSuffix(notices[index].Text, "\n") { + t.Fatalf("notice text was not normalized: %q", notices[index].Text) + } + if len(notices[index].SHA256) != 64 { + t.Fatalf("SHA-256 = %q", notices[index].SHA256) + } + } +} + +func TestReplacementLabelIsCheckoutIndependent(t *testing.T) { + root := filepath.Join(t.TempDir(), "repo") + inside := filepath.Join(root, "third_party", "module") + outside := filepath.Join(t.TempDir(), "module") + + if got := replacementLabel(root, &listedModule{Path: inside}); got != "./third_party/module" { + t.Fatalf("inside replacement = %q", got) + } + if got := replacementLabel(root, &listedModule{Path: outside}); got != "local replacement" { + t.Fatalf("outside replacement = %q", got) + } + if got := replacementLabel(root, &listedModule{Path: "example.com/fork", Version: "v1.2.3"}); got != "example.com/fork v1.2.3" { + t.Fatalf("module replacement = %q", got) + } +} + +func TestRenderUsesSafeMarkdownFenceAndProvenance(t *testing.T) { + generated, err := render([]dependency{{ + Path: "example.com/dependency", + Version: "v1.0.0", + Replacement: "./third_party/dependency", + Notices: []noticeFile{{ + Name: "LICENSE", + Text: "contains ``` fence\n", + SHA256: strings.Repeat("a", 64), + }}, + }}) + if err != nil { + t.Fatal(err) + } + text := string(generated) + for _, required := range []string{"````text", "example.com/dependency", "./third_party/dependency", "Brutal code provenance"} { + if !strings.Contains(text, required) { + t.Fatalf("generated notice does not contain %q", required) + } + } +} + +func TestReadNoticesRejectsMissingLicense(t *testing.T) { + if _, err := readNotices(t.TempDir()); err == nil { + t.Fatal("module without a root license was accepted") + } +} diff --git a/tools/vulnfilter/main.go b/tools/vulnfilter/main.go new file mode 100644 index 0000000..5f46952 --- /dev/null +++ b/tools/vulnfilter/main.go @@ -0,0 +1,140 @@ +// Command vulnfilter validates govulncheck's streaming JSON output. +// +// It fails on every reachable vulnerability except GO-2026-5288 when the +// finding is for the pinned Hysteria core v2.12.1. The upstream reviewed GHSA +// marks only versions <= 2.8.1 affected, while the automatically generated Go +// report currently says that every version is affected: +// +// https://github.com/advisories/GHSA-9fw6-xgg2-mq9q +// https://pkg.go.dev/vuln/GO-2026-5288 +// +// scripts/govulncheck.sh adds a source guard that also forbids enabling the +// vulnerable sniff feature while this narrowly scoped exception exists. +package main + +import ( + "encoding/json" + "errors" + "fmt" + "io" + "os" + "sort" + "strings" +) + +const ( + mainModule = "github.com/cppla/autocar" + allowedAdvisory = "GO-2026-5288" + allowedModule = "github.com/apernet/hysteria/core/v2" + allowedModuleVersion = "v2.12.1" +) + +type event struct { + Config *json.RawMessage `json:"config,omitempty"` + SBOM *json.RawMessage `json:"SBOM,omitempty"` + Finding *finding `json:"finding,omitempty"` +} + +type finding struct { + OSV string `json:"osv"` + Trace []frame `json:"trace"` +} + +type frame struct { + Module string `json:"module"` + Version string `json:"version"` + Package string `json:"package"` + Function string `json:"function"` +} + +func main() { + if err := validate(os.Stdin, os.Stdout); err != nil { + fmt.Fprintln(os.Stderr, "vulnerability scan:", err) + os.Exit(1) + } +} + +func validate(input io.Reader, output io.Writer) error { + decoder := json.NewDecoder(input) + seenConfig := false + seenSBOM := false + reachable := make(map[string][]finding) + moduleOnly := make(map[string]struct{}) + for { + var item event + err := decoder.Decode(&item) + if errors.Is(err, io.EOF) { + break + } + if err != nil { + return fmt.Errorf("decode govulncheck JSON: %w", err) + } + seenConfig = seenConfig || item.Config != nil + seenSBOM = seenSBOM || item.SBOM != nil + if item.Finding == nil || item.Finding.OSV == "" { + continue + } + if isReachable(*item.Finding) { + reachable[item.Finding.OSV] = append(reachable[item.Finding.OSV], *item.Finding) + } else { + moduleOnly[item.Finding.OSV] = struct{}{} + } + } + if !seenConfig || !seenSBOM { + return errors.New("incomplete govulncheck stream (missing config or SBOM)") + } + + delete(moduleOnly, allowedAdvisory) + if len(moduleOnly) != 0 { + fmt.Fprintf(output, "govulncheck: non-reachable module advisories: %s\n", strings.Join(sortedKeys(moduleOnly), ", ")) + } + + var unexpected []string + for id, findings := range reachable { + if id == allowedAdvisory && exactAllowedVersion(findings) { + fmt.Fprintf(output, "govulncheck: %s suppressed only for %s %s (upstream GHSA affects <= 2.8.1; sniff disabled)\n", id, allowedModule, allowedModuleVersion) + continue + } + unexpected = append(unexpected, id) + } + if len(unexpected) != 0 { + sort.Strings(unexpected) + return fmt.Errorf("reachable vulnerabilities: %s", strings.Join(unexpected, ", ")) + } + fmt.Fprintln(output, "govulncheck: no unsuppressed reachable vulnerabilities") + return nil +} + +func isReachable(item finding) bool { + for _, step := range item.Trace { + if step.Module == mainModule && step.Package != "" && step.Function != "" { + return true + } + } + return false +} + +func exactAllowedVersion(findings []finding) bool { + foundModule := false + for _, item := range findings { + for _, step := range item.Trace { + if step.Module != allowedModule { + continue + } + foundModule = true + if step.Version != allowedModuleVersion { + return false + } + } + } + return foundModule +} + +func sortedKeys(values map[string]struct{}) []string { + keys := make([]string, 0, len(values)) + for key := range values { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} diff --git a/tools/vulnfilter/main_test.go b/tools/vulnfilter/main_test.go new file mode 100644 index 0000000..27a9392 --- /dev/null +++ b/tools/vulnfilter/main_test.go @@ -0,0 +1,43 @@ +package main + +import ( + "bytes" + "strings" + "testing" +) + +func TestValidateAllowsOnlyPinnedFalsePositive(t *testing.T) { + stream := `{"config":{}} +{"SBOM":{}} +{"finding":{"osv":"GO-2026-5288","trace":[{"module":"github.com/apernet/hysteria/core/v2","version":"v2.12.1","package":"github.com/apernet/hysteria/core/v2/server","function":"NewServer"},{"module":"github.com/cppla/autocar","package":"github.com/cppla/autocar/internal/hy2","function":"Listen"}]}} +{"finding":{"osv":"GO-MODULE-ONLY","trace":[{"module":"example.invalid/dependency","version":"v1.0.0"}]}}` + var output bytes.Buffer + if err := validate(strings.NewReader(stream), &output); err != nil { + t.Fatal(err) + } + if !strings.Contains(output.String(), allowedAdvisory) || !strings.Contains(output.String(), "GO-MODULE-ONLY") { + t.Fatalf("output = %q", output.String()) + } +} + +func TestValidateRejectsUnexpectedReachableFinding(t *testing.T) { + stream := `{"config":{}}{"SBOM":{}}{"finding":{"osv":"GO-NEW","trace":[{"module":"bad.example/mod","version":"v1.0.0","package":"bad.example/mod/p","function":"Bad"},{"module":"github.com/cppla/autocar","package":"github.com/cppla/autocar/cmd/autocar","function":"main"}]}}` + if err := validate(strings.NewReader(stream), &bytes.Buffer{}); err == nil || !strings.Contains(err.Error(), "GO-NEW") { + t.Fatalf("error = %v", err) + } +} + +func TestValidateRejectsAllowedIDAtAnotherVersion(t *testing.T) { + stream := `{"config":{}}{"SBOM":{}}{"finding":{"osv":"GO-2026-5288","trace":[{"module":"github.com/apernet/hysteria/core/v2","version":"v2.8.1","package":"github.com/apernet/hysteria/core/v2/server","function":"NewServer"},{"module":"github.com/cppla/autocar","package":"github.com/cppla/autocar/internal/hy2","function":"Listen"}]}}` + if err := validate(strings.NewReader(stream), &bytes.Buffer{}); err == nil { + t.Fatal("affected Hysteria version was suppressed") + } +} + +func TestValidateRejectsIncompleteOrMalformedStream(t *testing.T) { + for _, stream := range []string{`{"config":{}}`, `{"config":{}} nope`} { + if err := validate(strings.NewReader(stream), &bytes.Buffer{}); err == nil { + t.Fatalf("stream %q was accepted", stream) + } + } +} From 66057c04404b759955db8bc736305996cfafc846 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Sat, 22 Aug 2026 14:07:34 +0000 Subject: [PATCH 2/2] build(deps): bump github.com/stretchr/testify Bumps the hardened-core-minor-and-patch group in /third_party/hysteria-core with 1 update: [github.com/stretchr/testify](https://github.com/stretchr/testify). Updates `github.com/stretchr/testify` from 1.11.1 to 1.12.1 - [Release notes](https://github.com/stretchr/testify/releases) - [Commits](https://github.com/stretchr/testify/compare/v1.11.1...v1.12.1) --- updated-dependencies: - dependency-name: github.com/stretchr/testify dependency-version: 1.12.1 dependency-type: direct:production update-type: version-update:semver-minor dependency-group: hardened-core-minor-and-patch ... Signed-off-by: dependabot[bot] --- third_party/hysteria-core/go.mod | 10 +++------- third_party/hysteria-core/go.sum | 26 ++++++-------------------- 2 files changed, 9 insertions(+), 27 deletions(-) diff --git a/third_party/hysteria-core/go.mod b/third_party/hysteria-core/go.mod index 98fbf11..5d1fca5 100644 --- a/third_party/hysteria-core/go.mod +++ b/third_party/hysteria-core/go.mod @@ -8,7 +8,7 @@ replace github.com/apernet/quic-go => ../quic-go require ( github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e - github.com/stretchr/testify v1.11.1 + github.com/stretchr/testify v1.12.1 go.uber.org/goleak v1.3.0 golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 golang.org/x/time v0.15.0 @@ -16,17 +16,13 @@ require ( require ( github.com/andybalholm/brotli v1.1.0 // indirect - github.com/davecgh/go-spew v1.1.1 // indirect github.com/klauspost/compress v1.18.7 // indirect - github.com/kr/text v0.2.0 // indirect - github.com/pmezard/go-difflib v1.0.0 // indirect github.com/quic-go/qpack v0.6.0 // indirect github.com/refraction-networking/utls v1.8.2 // indirect - github.com/rogpeppe/go-internal v1.12.0 // indirect - github.com/stretchr/objx v0.5.2 // indirect + github.com/stretchr/objx v0.5.3 // indirect + go.yaml.in/yaml/v3 v3.0.5 // indirect golang.org/x/crypto v0.54.0 // indirect golang.org/x/net v0.57.0 // indirect golang.org/x/sys v0.47.0 // indirect golang.org/x/text v0.40.0 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/third_party/hysteria-core/go.sum b/third_party/hysteria-core/go.sum index effebfa..679b980 100644 --- a/third_party/hysteria-core/go.sum +++ b/third_party/hysteria-core/go.sum @@ -1,32 +1,23 @@ github.com/andybalholm/brotli v1.1.0 h1:eLKJA0d02Lf0mVpIDgYnqXcUn0GqVmEFny3VuID1U3M= github.com/andybalholm/brotli v1.1.0/go.mod h1:sms7XGricyQI9K10gOSf56VKKWS4oLer58Q+mhRPtnY= -github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= -github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= -github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= -github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= -github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= -github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= -github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= -github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= -github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0= github.com/quic-go/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk= github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEvV+S9iJ2IdQo= github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= -github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= -github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= -github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= -github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= -github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= -github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/stretchr/objx v0.5.3 h1:jmXUvGomnU1o3W/V5h2VEradbpJDwGrzugQQvL0POH4= +github.com/stretchr/objx v0.5.3/go.mod h1:rDQraq+vQZU7Fde9LOZLr8Tax6zZvy4kuNKF+QYS+U0= +github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= +github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko= go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o= +go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= +go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM= @@ -39,8 +30,3 @@ golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= -gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= -gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= -gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=